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")
logger.info(f"根据archive_message_id={archive_message_id},查询结果为: {media_file}")
if not media_file:
return ASRRecognizeResponse(
success=False,
@ -73,15 +74,16 @@ class ASRService:
try:
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:
return ASRRecognizeResponse(
success=False,
message="ASR 识别结果为空",
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)
logger.info(f"更新 archive_message_id={archive_message_id} 的结果为: {update_ok}")
if update_ok:
await self._mark_processed(archive_message_id)
return ASRRecognizeResponse(

View File

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

View File

@ -2,12 +2,16 @@ from typing import Annotated
from fastapi import APIRouter
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.service.query_qa_service import QueryQAService
from app.client.embedding_client_manager import embedding_client
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()
@ -17,3 +21,26 @@ async def query_handler(query: QueryQASchema, query_service: Annotated[QueryQASe
logger.info(f"请求参数:{query.external_str},历史top15:{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)
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,
qa_milvus_repository=qa_milvus_repository
)

View File

@ -1,4 +1,5 @@
from typing import List, Optional
from datetime import datetime, timedelta
from sqlalchemy import text, update
from sqlalchemy.ext.asyncio import AsyncSession
@ -93,4 +94,65 @@ class ArchiveMessagesRepository:
return result.rowcount > 0
except Exception as e:
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 json
import time
from datetime import datetime, timedelta
from typing import List
from langchain_core.output_parsers import StrOutputParser
from langchain_core.prompts import PromptTemplate
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})
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):
logger.info(f"执行日期:{day}")
@ -211,7 +238,7 @@ async def build(day:str):
archive_messages_repository = ArchiveMessagesRepository(db_assistant)
summary_repository = SummaryRepository()
service = SummaryService(archive_messages_repository, embedding_client, summary_repository)
await service.get_message_summary(day)
await service.real_time_summary("BuShouShiJinBuGaiMing","wmI1AkDQAADMN3I0RScMP_g4R8jGoVfQ")
finally:
await embedding_client.close()
await db_assistant_mysql_client_manager.close()