feat: 意图识别优化
This commit is contained in:
parent
ac96ab14ea
commit
928085ecfc
@ -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(
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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
|
||||
)
|
||||
|
||||
@ -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()]
|
||||
|
||||
@ -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()
|
||||
|
||||
Loading…
Reference in New Issue
Block a user