feat: 意图识别优化

This commit is contained in:
qinyong@9artedu.com 2026-08-10 18:02:08 +08:00
parent ac96ab14ea
commit 928085ecfc
6 changed files with 127 additions and 9 deletions

View File

@ -57,6 +57,7 @@ class ASRService:
) )
media_file = await self.media_repository.get_by_archive_message_id(archive_message_id, file_type="voice") media_file = await self.media_repository.get_by_archive_message_id(archive_message_id, file_type="voice")
logger.info(f"根据archive_message_id={archive_message_id},查询结果为: {media_file}")
if not media_file: if not media_file:
return ASRRecognizeResponse( return ASRRecognizeResponse(
success=False, success=False,
@ -73,15 +74,16 @@ class ASRService:
try: try:
result = await asr_service.recognize_from_url(media_file.cos_url) result = await asr_service.recognize_from_url(media_file.cos_url)
logger.info(f"archive_message_id={archive_message_id},ASR 最终识别结果: {result}")
if not result: if not result:
return ASRRecognizeResponse( return ASRRecognizeResponse(
success=False, success=False,
message="ASR 识别结果为空", message="ASR 识别结果为空",
archive_message_id=archive_message_id archive_message_id=archive_message_id
) )
logger.info(f"更新 archive_message_id={archive_message_id} 的 content 为: {result}")
update_ok = await self.msg_repository.update_content_by_id(archive_message_id, result) update_ok = await self.msg_repository.update_content_by_id(archive_message_id, result)
logger.info(f"更新 archive_message_id={archive_message_id} 的结果为: {update_ok}")
if update_ok: if update_ok:
await self._mark_processed(archive_message_id) await self._mark_processed(archive_message_id)
return ASRRecognizeResponse( return ASRRecognizeResponse(

View File

@ -473,8 +473,8 @@ async def build(day:str):
if __name__ == '__main__': if __name__ == '__main__':
start = datetime(2026, 7, 29) start = datetime(2026, 8, 9)
end = datetime(2026, 7, 29) end = datetime(2026, 8, 9)
result = [] result = []
temp = start temp = start

View File

@ -2,12 +2,16 @@ from typing import Annotated
from fastapi import APIRouter from fastapi import APIRouter
from fastapi.params import Depends from fastapi.params import Depends
from watchfiles import awatch from sqlalchemy.ext.asyncio import AsyncSession
from app.chat_qa_query.api.qa_dependencies import get_qa_query_service from app.chat_qa_query.api.qa_dependencies import get_assistant_session, get_qa_query_service
from app.chat_qa_query.api.schema.query_qa_schema import QueryQASchema from app.chat_qa_query.api.schema.query_qa_schema import QueryQASchema
from app.chat_qa_query.service.query_qa_service import QueryQAService from app.chat_qa_query.service.query_qa_service import QueryQAService
from app.client.embedding_client_manager import embedding_client
from app.core.log import logger from app.core.log import logger
from app.repository.archive_messages_repository import ArchiveMessagesRepository
from app.repository.milvus.summary_repository import SummaryRepository
from app.summary.service import SummaryService
query_qa_router = APIRouter() query_qa_router = APIRouter()
@ -17,3 +21,26 @@ async def query_handler(query: QueryQASchema, query_service: Annotated[QueryQASe
logger.info(f"请求参数:{query.external_str},历史top15:{query.messages_top15}") logger.info(f"请求参数:{query.external_str},历史top15:{query.messages_top15}")
return await query_service.query(query.external_str, query.messages_top15) return await query_service.query(query.external_str, query.messages_top15)
async def get_summary_service(
session: Annotated[AsyncSession, Depends(get_assistant_session)],
) -> SummaryService:
archive_messages_repository = ArchiveMessagesRepository(session)
summary_repository = SummaryRepository()
return SummaryService(
archive_messages_repository=archive_messages_repository,
embedding_client=embedding_client,
summary_repository=summary_repository,
)
@query_qa_router.get("/api/summary/real-time")
async def real_time_summary_handler(
from_user: str,
to_user: str,
summary_service: Annotated[SummaryService, Depends(get_summary_service)],
):
logger.info(f"实时总结请求: from_user={from_user}, to_user={to_user}")
result = await summary_service.real_time_summary(from_user, to_user)
return result

View File

@ -45,7 +45,7 @@ if __name__ == '__main__':
archive_messages_mysql_repository = ArchiveMessagesRepository(db_session) archive_messages_mysql_repository = ArchiveMessagesRepository(db_session)
qa_milvus_repository = QARepository() qa_milvus_repository = QARepository()
state: QueryQAState = QueryQAState(original_query="报名有没有优惠,另外 如果和同学一起报名 有没有什么优惠", history=[]) state: QueryQAState = QueryQAState(original_query="老师,我想问一下,是不是学费涨了", history=[])
context = QueryQAContext(meta_mysql_repository=archive_messages_mysql_repository, context = QueryQAContext(meta_mysql_repository=archive_messages_mysql_repository,
qa_milvus_repository=qa_milvus_repository qa_milvus_repository=qa_milvus_repository
) )

View File

@ -1,4 +1,5 @@
from typing import List, Optional from typing import List, Optional
from datetime import datetime, timedelta
from sqlalchemy import text, update from sqlalchemy import text, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@ -94,3 +95,64 @@ class ArchiveMessagesRepository:
except Exception as e: except Exception as e:
await self.session.rollback() await self.session.rollback()
raise e raise e
async def get_conversation_messages(
self,
from_user: str,
to_user: str,
start_time: Optional[int] = None,
end_time: Optional[int] = None,
) -> List[ArchiveMessages]:
toUser_m = "[\"" + to_user + "\"]"
fromUser_m = "[\"" + from_user + "\"]"
# 默认时间范围:当天开始(含) ~ 次日开始(不含)单位UTC 毫秒
if start_time is None or end_time is None:
today_start = datetime.now().replace(hour=0, minute=0, second=0, microsecond=0)
next_day_start = today_start + timedelta(days=1)
if start_time is None:
start_time = int(today_start.timestamp() * 1000)
if end_time is None:
end_time = int(next_day_start.timestamp() * 1000)
sql = """
(
SELECT msgtime, created_at, from_role, content
FROM archive_messages
WHERE from_user = :from_user_ab
AND to_user = :to_user_ab
AND msgtime >= :start_time
AND msgtime < :end_time
AND msgtype IN ('text', 'voice')
AND content IS NOT NULL
AND content != ''
LIMIT 100
)
UNION ALL
(
SELECT msgtime, created_at, from_role, content
FROM archive_messages
WHERE from_user = :from_user_ba
AND to_user = :to_user_ba
AND msgtime >= :start_time
AND msgtime < :end_time
AND msgtype IN ('text', 'voice')
AND content IS NOT NULL
AND content != ''
LIMIT 100
)
ORDER BY msgtime ASC
"""
params = {
# 方向1: from_user -> to_user
"from_user_ab": from_user,
"to_user_ab": toUser_m,
# 方向2: to_user -> from_user
"from_user_ba": to_user,
"to_user_ba": fromUser_m,
# 共用时间范围
"start_time": start_time,
"end_time": end_time,
}
result = await self.session.execute(text(sql), params)
return [ArchiveMessages(**dict(row)) for row in result.mappings().fetchall()]

View File

@ -1,8 +1,10 @@
import asyncio import asyncio
import json
import time import time
from datetime import datetime, timedelta from datetime import datetime, timedelta
from typing import List from typing import List
from langchain_core.output_parsers import StrOutputParser
from langchain_core.prompts import PromptTemplate from langchain_core.prompts import PromptTemplate
from app.client.embedding_client_manager import EmbeddingClientManager, embedding_client from app.client.embedding_client_manager import EmbeddingClientManager, embedding_client
@ -200,6 +202,31 @@ class SummaryService:
result = await chain.ainvoke({"message_str": message_str,"day_date":day_date}) result = await chain.ainvoke({"message_str": message_str,"day_date":day_date})
return result return result
# 实时总结
async def real_time_summary(self, from_user: str, to_user: str) -> str:
message_list = await self.archive_messages_repository.get_conversation_messages(from_user, to_user)
if not message_list:
return ""
prompt = """
# 聊天历史轻量化摘要任务
任务生成极简轮次摘要用于替换原始对话节省上下文窗口
# 历史对话
{message_str}
## 说明
其中INTERNAL表示老师EXTERNAL表示学生
输出格式示例无多余文字
学生反馈订单 #12345 物流停滞超7天未更新。
询问是否支持跨店满减叠加优惠券
"""
lines = [f" {msg.created_at}:{msg.from_role}{msg.content}" for msg in message_list]
message_str = "\n".join(lines)
prompt_template = PromptTemplate(template=prompt, input_variables=["message_str"])
chain = prompt_template | llm | StrOutputParser()
result = await chain.ainvoke({"message_str": message_str})
logger.info(f"{from_user},{to_user}实时总结结果:{result}")
return result
async def build(day:str): async def build(day:str):
logger.info(f"执行日期:{day}") logger.info(f"执行日期:{day}")
@ -211,7 +238,7 @@ async def build(day:str):
archive_messages_repository = ArchiveMessagesRepository(db_assistant) archive_messages_repository = ArchiveMessagesRepository(db_assistant)
summary_repository = SummaryRepository() summary_repository = SummaryRepository()
service = SummaryService(archive_messages_repository, embedding_client, summary_repository) service = SummaryService(archive_messages_repository, embedding_client, summary_repository)
await service.get_message_summary(day) await service.real_time_summary("BuShouShiJinBuGaiMing","wmI1AkDQAADMN3I0RScMP_g4R8jGoVfQ")
finally: finally:
await embedding_client.close() await embedding_client.close()
await db_assistant_mysql_client_manager.close() await db_assistant_mysql_client_manager.close()