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)