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")
|
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(
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|||||||
@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@ -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()]
|
||||||
|
|||||||
@ -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()
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user