sales-assistant-py-new/app/asr/api/asr_dependencies.py

112 lines
4.6 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")
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)
if not result:
return ASRRecognizeResponse(
success=False,
message="ASR 识别结果为空",
archive_message_id=archive_message_id
)
update_ok = await self.msg_repository.update_content_by_id(archive_message_id, result)
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)