114 lines
5.0 KiB
Python
114 lines
5.0 KiB
Python
from typing import Annotated
|
|
|
|
from fastapi import Depends
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.asr.asr_service import asr_service
|
|
from app.asr.api.schema.asr_schema import ASRRecognizeResponse
|
|
from app.client.mysql_client_manager import db_assistant_mysql_client_manager
|
|
from app.client.redis_client_manager import redis_client_manager
|
|
from app.repository.archive_media_files_repository import ArchiveMediaFilesRepository
|
|
from app.repository.archive_messages_repository import ArchiveMessagesRepository
|
|
from app.core.log import logger
|
|
|
|
|
|
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_media_files_repository(session: Annotated[AsyncSession, Depends(get_assistant_session)]) -> ArchiveMediaFilesRepository:
|
|
return ArchiveMediaFilesRepository(session)
|
|
|
|
|
|
async def get_archive_messages_repository(session: Annotated[AsyncSession, Depends(get_assistant_session)]) -> ArchiveMessagesRepository:
|
|
return ArchiveMessagesRepository(session)
|
|
|
|
|
|
class ASRService:
|
|
def __init__(self,
|
|
media_repository: Annotated[ArchiveMediaFilesRepository, Depends(get_archive_media_files_repository)],
|
|
msg_repository: Annotated[ArchiveMessagesRepository, Depends(get_archive_messages_repository)]):
|
|
self.media_repository = media_repository
|
|
self.msg_repository = msg_repository
|
|
|
|
def _get_redis_key(self, message_id: int) -> str:
|
|
return f"archive_messages_{message_id}"
|
|
|
|
async def _is_processed(self, message_id: int) -> bool:
|
|
try:
|
|
return await redis_client_manager.client.exists(self._get_redis_key(message_id)) > 0
|
|
except Exception as e:
|
|
logger.warning(f"Redis 检查失败: {e}")
|
|
return False
|
|
|
|
async def _mark_processed(self, message_id: int):
|
|
try:
|
|
await redis_client_manager.client.set(self._get_redis_key(message_id), "1", ex=604800)
|
|
except Exception as e:
|
|
logger.warning(f"Redis 标记失败: {e}")
|
|
|
|
async def recognize(self, archive_message_id: int) -> ASRRecognizeResponse:
|
|
if await self._is_processed(archive_message_id):
|
|
return ASRRecognizeResponse(
|
|
success=False,
|
|
message=f"archive_message_id={archive_message_id} 已处理过",
|
|
archive_message_id=archive_message_id
|
|
)
|
|
|
|
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,
|
|
message=f"未找到 archive_message_id={archive_message_id} 对应的语音文件",
|
|
archive_message_id=archive_message_id
|
|
)
|
|
|
|
if not media_file.cos_url:
|
|
return ASRRecognizeResponse(
|
|
success=False,
|
|
message=f"媒体文件不存在 cos_url",
|
|
archive_message_id=archive_message_id
|
|
)
|
|
|
|
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(
|
|
success=True,
|
|
message="ASR 识别成功并更新到数据库",
|
|
archive_message_id=archive_message_id,
|
|
content=result
|
|
)
|
|
else:
|
|
return ASRRecognizeResponse(
|
|
success=False,
|
|
message="ASR 识别成功但更新数据库失败",
|
|
archive_message_id=archive_message_id,
|
|
content=result
|
|
)
|
|
except Exception as e:
|
|
return ASRRecognizeResponse(
|
|
success=False,
|
|
message=f"ASR 识别失败: {str(e)}",
|
|
archive_message_id=archive_message_id
|
|
)
|
|
|
|
|
|
async def get_asr_service(
|
|
media_repository: Annotated[ArchiveMediaFilesRepository, Depends(get_archive_media_files_repository)],
|
|
msg_repository: Annotated[ArchiveMessagesRepository, Depends(get_archive_messages_repository)]
|
|
) -> ASRService:
|
|
return ASRService(media_repository=media_repository, msg_repository=msg_repository) |