sales-assistant-py-new/app/asr/asr_voice_processor.py

130 lines
5.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import asyncio
import logging
from app.asr.asr_service import asr_service
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
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
class ASRVoiceProcessor:
def __init__(self):
db_assistant_mysql_client_manager.init()
redis_client_manager.init()
self.session_factory = db_assistant_mysql_client_manager.session_factory
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 process_voice_files(self, batch_size: int = 10, max_records: int = None):
"""
处理历史语音文件
Args:
batch_size: 每批处理数量
max_records: 最大处理记录数None表示处理所有
"""
async with self.session_factory() as session:
repository = ArchiveMediaFilesRepository(session)
msg_repository = ArchiveMessagesRepository(session)
total = await repository.count_voice_files()
logger.info(f"语音文件总数: {total}")
processed = 0
success_count = 0
skip_count = 0
offset = 0
while True:
if max_records and processed >= max_records:
logger.info(f"已达到最大处理记录数: {max_records}")
break
voice_files = await repository.get_voice_files(limit=batch_size, offset=offset)
logger.info(f"处理语音文件: 数量={len(voice_files)}")
if not voice_files:
logger.info("没有更多语音文件需要处理")
break
logger.info(f"处理批次: 偏移={offset}, 数量={len(voice_files)}")
for i,voice_file in enumerate(voice_files):
logger.info(f"开始处理第{i} 音频文件, 总数={len(voice_files)}")
if not voice_file.cos_url:
processed += 1
continue
archive_message_id = voice_file.archive_message_id
if not archive_message_id:
logger.warning(f"记录 ID: {voice_file.id} 缺少 archive_message_id跳过")
processed += 1
continue
if await self._is_processed(archive_message_id):
logger.info(f"记录 ID: {voice_file.id}, archive_message_id: {archive_message_id} 已处理过,跳过")
skip_count += 1
processed += 1
continue
try:
logger.info(f"处理记录 ID: {voice_file.id}, archive_message_id: {archive_message_id}, URL: {voice_file.cos_url}")
result = await asr_service.recognize_from_url(voice_file.cos_url)
logger.info(f"识别结果: {result}")
if result:
update_ok = await msg_repository.update_content_by_id(archive_message_id, result)
if update_ok:
await self._mark_processed(archive_message_id)
success_count += 1
logger.info(f"已更新成功: archive_message_id={archive_message_id}")
else:
logger.warning(f"更新失败: archive_message_id={archive_message_id}")
except Exception as e:
logger.error(f"处理失败 ID: {voice_file.id}, 错误: {str(e)}")
processed += 1
if max_records and processed >= max_records:
break
offset += batch_size
await asyncio.sleep(0.5)
logger.info(f"处理完成,总计处理: {processed} 条记录,成功: {success_count} 条,跳过: {skip_count}")
async def close(self):
await db_assistant_mysql_client_manager.close()
await redis_client_manager.close()
async def main():
processor = ASRVoiceProcessor()
try:
await processor.process_voice_files(batch_size=10, max_records=10000)
finally:
await processor.close()
if __name__ == "__main__":
asyncio.run(main())