From 699aea73a699d0f83212880c72326520493116b3 Mon Sep 17 00:00:00 2001 From: "qinyong@9artedu.com" Date: Fri, 26 Jun 2026 10:00:18 +0800 Subject: [PATCH] =?UTF-8?q?=E5=A2=9E=E5=8A=A0asr=E8=83=BD=E5=8A=9B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- app/asr/__init__.py | 4 + app/asr/api/asr_dependencies.py | 112 +++++++ app/asr/api/router/asr_router.py | 15 + app/asr/api/schema/asr_schema.py | 12 + app/asr/asr_client.py | 137 +++++++++ app/asr/asr_service.py | 61 ++++ app/asr/asr_voice_processor.py | 129 ++++++++ app/client/redis_client_manager.py | 71 +++++ app/conf/app_config.py | 16 + app/core/liftspan.py | 30 +- app/models/mysql/__init__.py | 3 +- .../archive_media_files_repository.py | 60 ++++ app/repository/archive_messages_repository.py | 36 ++- app/repository/redis/__init__.py | 0 app/repository/redis/redis_repository.py | 283 ++++++++++++++++++ app/summary/service.py | 10 +- app/uvicorn_main.py | 2 + conf/app_config.yaml | 12 + pyproject.toml | 1 + uv.lock | 11 + 20 files changed, 989 insertions(+), 16 deletions(-) create mode 100644 app/asr/__init__.py create mode 100644 app/asr/api/asr_dependencies.py create mode 100644 app/asr/api/router/asr_router.py create mode 100644 app/asr/api/schema/asr_schema.py create mode 100644 app/asr/asr_client.py create mode 100644 app/asr/asr_service.py create mode 100644 app/asr/asr_voice_processor.py create mode 100644 app/client/redis_client_manager.py create mode 100644 app/repository/archive_media_files_repository.py create mode 100644 app/repository/redis/__init__.py create mode 100644 app/repository/redis/redis_repository.py diff --git a/app/asr/__init__.py b/app/asr/__init__.py new file mode 100644 index 0000000..866139c --- /dev/null +++ b/app/asr/__init__.py @@ -0,0 +1,4 @@ +from .asr_client import Qwen3ASRClient, asr_client +from .asr_service import ASRService, asr_service + +__all__ = ["Qwen3ASRClient", "asr_client", "ASRService", "asr_service"] \ No newline at end of file diff --git a/app/asr/api/asr_dependencies.py b/app/asr/api/asr_dependencies.py new file mode 100644 index 0000000..af0ef85 --- /dev/null +++ b/app/asr/api/asr_dependencies.py @@ -0,0 +1,112 @@ +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) \ No newline at end of file diff --git a/app/asr/api/router/asr_router.py b/app/asr/api/router/asr_router.py new file mode 100644 index 0000000..c51f66d --- /dev/null +++ b/app/asr/api/router/asr_router.py @@ -0,0 +1,15 @@ +from typing import Annotated + +from fastapi import APIRouter, Depends + +from app.asr.api.asr_dependencies import ASRService, get_asr_service +from app.asr.api.schema.asr_schema import ASRRecognizeRequest, ASRRecognizeResponse +from app.core.log import logger + +asr_router = APIRouter() + + +@asr_router.post("/api/asr/recognize", response_model=ASRRecognizeResponse) +async def recognize_handler(request: ASRRecognizeRequest, asr_service: Annotated[ASRService, Depends(get_asr_service)]): + logger.info(f"ASR 识别请求: archive_message_id={request.archive_message_id}") + return await asr_service.recognize(request.archive_message_id) \ No newline at end of file diff --git a/app/asr/api/schema/asr_schema.py b/app/asr/api/schema/asr_schema.py new file mode 100644 index 0000000..2cabb11 --- /dev/null +++ b/app/asr/api/schema/asr_schema.py @@ -0,0 +1,12 @@ +from pydantic import BaseModel + + +class ASRRecognizeRequest(BaseModel): + archive_message_id: int + + +class ASRRecognizeResponse(BaseModel): + success: bool + message: str + archive_message_id: int + content: str = None \ No newline at end of file diff --git a/app/asr/asr_client.py b/app/asr/asr_client.py new file mode 100644 index 0000000..f35d7d2 --- /dev/null +++ b/app/asr/asr_client.py @@ -0,0 +1,137 @@ +import asyncio +import logging +from typing import Optional + +import httpx + +from app.conf.app_config import ASRConfig, app_config + +logger = logging.getLogger(__name__) + + +class Qwen3ASRClient: + def __init__(self, config: ASRConfig): + self.config = config + self._session: Optional[httpx.AsyncClient] = None + + @property + def session(self) -> httpx.AsyncClient: + if self._session is None: + timeout = httpx.Timeout(120.0, connect=10.0, read=60.0) + headers = {} + if self.config.api_key and self.config.api_key.strip(): + headers["Authorization"] = f"Bearer {self.config.api_key.strip()}" + self._session = httpx.AsyncClient( + timeout=timeout, + follow_redirects=True, + headers=headers + ) + return self._session + + def _get_submit_url(self) -> str: + return f"{self.config.base_url}/api/v1/services/audio/asr/transcription" + + def _get_task_url(self, task_id: str) -> str: + return f"{self.config.base_url}/api/v1/tasks/{task_id}" + + async def _submit_task(self, audio_url: str) -> str: + url = self._get_submit_url() + + payload = { + "model": "qwen3-asr-flash-filetrans", + "input": { + "file_url": audio_url + }, + "parameters": { + "enable_itn": False + } + } + + headers = { + "Content-Type": "application/json", + "X-DashScope-Async": "enable" + } + + response = await self.session.post(url, json=payload, headers=headers) + if response.status_code >= 400: + logger.error(f"ASR任务提交失败,状态码: {response.status_code}, 响应: {response.text}") + response.raise_for_status() + result = response.json() + + if "output" in result and "task_id" in result["output"]: + task_id = result["output"]["task_id"] + logger.info(f"ASR任务提交成功,task_id: {task_id}") + return task_id + else: + logger.error(f"ASR任务提交失败,响应: {result}") + raise ValueError(f"Failed to submit ASR task: {result}") + + async def _poll_task_result(self, task_id: str, max_attempts: int = 60, interval: float = 2.0) -> str: + url = self._get_task_url(task_id) + + for attempt in range(max_attempts): + response = await self.session.get(url) + response.raise_for_status() + result = response.json() + + task_status = result.get("output", {}).get("task_status", "") + + if task_status == "SUCCEEDED": + logger.info(f"ASR任务完成,task_id: {task_id}") + output = result.get("output", {}) + result_data = output.get("result", {}) + transcription_url = result_data.get("transcription_url", "") + + if transcription_url: + logger.info(f"下载识别结果: {transcription_url}") + trans_response = await self.session.get(transcription_url) + trans_response.raise_for_status() + trans_data = trans_response.json() + + texts = [] + if "transcripts" in trans_data: + for item in trans_data["transcripts"]: + if "text" in item: + texts.append(item["text"]) + elif "text" in trans_data: + texts.append(trans_data["text"]) + + text = "\n".join(texts).strip() + else: + text = str(output) + + logger.info(f"识别结果长度: {len(text)}") + return text + elif task_status in ("FAILED", "UNKNOWN"): + error_msg = result.get("output", {}).get("message", f"Task {task_status}") + logger.error(f"ASR任务失败,task_id: {task_id}, 错误: {error_msg}") + raise RuntimeError(f"ASR task failed: {error_msg}") + + logger.debug(f"ASR任务进行中,task_id: {task_id}, 状态: {task_status}, 第{attempt + 1}次轮询") + await asyncio.sleep(interval) + + raise TimeoutError(f"ASR task timed out after {max_attempts} attempts") + + async def recognize_from_url_direct(self, audio_url: str) -> str: + logger.info(f"开始ASR URL直传识别,URL: {audio_url}") + try: + task_id = await self._submit_task(audio_url) + result = await self._poll_task_result(task_id) + return result + except Exception as e: + logger.error(f"ASR URL直传识别失败: {str(e)}", exc_info=True) + raise + + async def recognize(self, audio_data: bytes, format: str = "wav", **kwargs) -> str: + raise NotImplementedError("Audio data upload not implemented yet, use recognize_from_url_direct") + + async def recognize_from_file(self, file_path: str) -> str: + raise NotImplementedError("Local file upload not implemented yet, use recognize_from_url_direct") + + async def close(self): + if self._session: + await self._session.aclose() + self._session = None + + +asr_client = Qwen3ASRClient(app_config.asr) \ No newline at end of file diff --git a/app/asr/asr_service.py b/app/asr/asr_service.py new file mode 100644 index 0000000..da8b07c --- /dev/null +++ b/app/asr/asr_service.py @@ -0,0 +1,61 @@ +import asyncio +import logging +from typing import Any, Dict + +from langchain_core.runnables import RunnableLambda + +from app.asr.asr_client import asr_client +from app.core.log import logger + + +class ASRService: + def __init__(self): + self._recognize_chain = RunnableLambda(self._recognize_handler) + self._recognize_file_chain = RunnableLambda(self._recognize_file_handler) + self._recognize_url_chain = RunnableLambda(self._recognize_url_handler) + + async def _recognize_handler(self, input_data: Dict[str, Any], **kwargs) -> str: + audio_data = input_data.get("audio_data") + format = input_data.get("format", "wav") + logger.info(f"开始语音识别,数据长度: {len(audio_data) if audio_data else 0}") + result = await asr_client.recognize(audio_data, format) + logger.info(f"语音识别完成,结果长度: {len(result)}") + return result + + async def _recognize_file_handler(self, input_data: Dict[str, Any], **kwargs) -> str: + file_path = input_data.get("file_path") + logger.info(f"开始语音识别,文件路径: {file_path}") + result = await asr_client.recognize_from_file(file_path) + logger.info(f"语音识别完成,结果长度: {len(result)}") + return result + + async def _recognize_url_handler(self, input_data: Dict[str, Any], **kwargs) -> str: + audio_url = input_data.get("audio_url") + logger.info(f"开始语音识别,URL: {audio_url}") + result = await asr_client.recognize_from_url_direct(audio_url) + logger.info(f"语音识别完成,结果长度: {len(result)}") + return result + + async def recognize(self, audio_data: bytes, format: str = "wav") -> str: + return await self._recognize_chain.ainvoke({"audio_data": audio_data, "format": format}) + + async def recognize_from_file(self, file_path: str) -> str: + return await self._recognize_file_chain.ainvoke({"file_path": file_path}) + + async def recognize_from_url(self, audio_url: str) -> str: + return await self._recognize_url_chain.ainvoke({"audio_url": audio_url}) + + async def close(self): + await asr_client.close() + + +asr_service = ASRService() + +if __name__ == "__main__": + async def test(): + # 测试用户的音频URL + audio_url = "https://meedu-cos.9artedu.com/archive/media/10780491256375083103_1782355020252_external.amr" + print(f"测试音频URL: {audio_url}") + result = await asr_service.recognize_from_url(audio_url) + print(f"识别结果: {result}") + asyncio.run(test()) \ No newline at end of file diff --git a/app/asr/asr_voice_processor.py b/app/asr/asr_voice_processor.py new file mode 100644 index 0000000..7b29aaa --- /dev/null +++ b/app/asr/asr_voice_processor.py @@ -0,0 +1,129 @@ +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()) diff --git a/app/client/redis_client_manager.py b/app/client/redis_client_manager.py new file mode 100644 index 0000000..5dd607a --- /dev/null +++ b/app/client/redis_client_manager.py @@ -0,0 +1,71 @@ +from typing import Optional + +import redis +from redis import asyncio as aioredis + +from app.conf.app_config import RedisConfig, app_config + + +class RedisClientManager: + def __init__(self, config: RedisConfig): + self.config = config + self._client: Optional[aioredis.Redis] = None + self._sync_client: Optional[redis.Redis] = None + + def init(self): + self._client = aioredis.from_url( + f"redis://{self.config.host}:{self.config.port}", + password=self.config.password if self.config.password else None, + db=self.config.db, + decode_responses=self.config.decode_responses, + socket_timeout=10, + socket_connect_timeout=10 + ) + self._sync_client = redis.Redis( + host=self.config.host, + port=self.config.port, + password=self.config.password if self.config.password else None, + db=self.config.db, + decode_responses=self.config.decode_responses, + socket_timeout=10, + socket_connect_timeout=10 + ) + + @property + def client(self) -> aioredis.Redis: + if self._client is None: + self.init() + return self._client + + @property + def sync_client(self) -> redis.Redis: + if self._sync_client is None: + self.init() + return self._sync_client + + async def close(self): + if self._client: + await self._client.close() + if self._sync_client: + self._sync_client.close() + + +redis_client_manager = RedisClientManager(app_config.redis) + +if __name__ == '__main__': + import asyncio + + redis_client_manager.init() + + async def test(): + await redis_client_manager.client.set("test_key", "test_value") + value = await redis_client_manager.client.get("test_key") + print(f"Async get: {value}") + + redis_client_manager.sync_client.set("sync_test_key", "sync_test_value") + sync_value = redis_client_manager.sync_client.get("sync_test_key") + print(f"Sync get: {sync_value}") + + await redis_client_manager.close() + + asyncio.run(test()) \ No newline at end of file diff --git a/app/conf/app_config.py b/app/conf/app_config.py index a316f2a..e0f6b47 100644 --- a/app/conf/app_config.py +++ b/app/conf/app_config.py @@ -54,6 +54,20 @@ class LLMConfig: api_key: str base_url: str +@dataclass +class ASRConfig: + model_name: str + api_key: str + base_url: str + +@dataclass +class RedisConfig: + host: str + port: int + password: str + db: int = 0 + decode_responses: bool = True + @dataclass class AppConfig: logging: LoggingConfig @@ -61,6 +75,8 @@ class AppConfig: embedding: EmbeddingConfig llm: LLMConfig milvus: MilvusConfig + asr: ASRConfig + redis: RedisConfig config_file = Path(__file__).parents[2] / 'conf' / 'app_config.yaml' diff --git a/app/core/liftspan.py b/app/core/liftspan.py index 619b0cc..7919743 100644 --- a/app/core/liftspan.py +++ b/app/core/liftspan.py @@ -7,6 +7,7 @@ from fastapi import FastAPI from app.client.embedding_client_manager import embedding_client from app.client.milvus_client_manager import milvus_client from app.client.mysql_client_manager import db_assistant_mysql_client_manager +from app.client.redis_client_manager import redis_client_manager from app.core.log import logger from app.summary.service import build @@ -21,24 +22,40 @@ scheduler = AsyncIOScheduler(timezone='Asia/Shanghai') async def generate_daily_summary(): yesterday = datetime.now() - timedelta(days=1) date = datetime(yesterday.year, yesterday.month, yesterday.day) - # 按 年-月-日 无补零格式输出 date_format = f"{date.year}-{date.month}-{date.day}" - #date_format = "2026-06-23" - await build(date_format) - logger.info(f"{date_format} 消息摘要生成完成") + lock_key = f"daily_summary_lock_{date_format}" + lock_expire = 60 + + try: + acquired = await redis_client_manager.client.set(lock_key, "1", ex=lock_expire, nx=True) + if not acquired: + logger.info(f"{date_format} 消息摘要任务已在其他实例执行,跳过") + return + + logger.info(f"{date_format} 消息摘要任务获取锁成功,开始执行") + await build(date_format) + logger.info(f"{date_format} 消息摘要生成完成") + except Exception as e: + logger.error(f"{date_format} 消息摘要任务执行失败: {str(e)}") + finally: + try: + await redis_client_manager.client.delete(lock_key) + except Exception as e: + logger.warning(f"释放锁失败: {e}") @asynccontextmanager async def lifespan(app: FastAPI): db_assistant_mysql_client_manager.init() embedding_client.init() milvus_client.init() + redis_client_manager.init() scheduler.add_job( generate_daily_summary, trigger='cron', - hour=9, - minute=40, + hour=2, + minute=0, second=0, id='daily_summary', name='每日消息摘要任务', @@ -50,5 +67,6 @@ async def lifespan(app: FastAPI): await db_assistant_mysql_client_manager.close() await embedding_client.close() milvus_client.close() + await redis_client_manager.close() scheduler.shutdown() logger.info("定时任务调度器已关闭") \ No newline at end of file diff --git a/app/models/mysql/__init__.py b/app/models/mysql/__init__.py index ca4cf59..75cf842 100644 --- a/app/models/mysql/__init__.py +++ b/app/models/mysql/__init__.py @@ -1,3 +1,4 @@ from .archive_messages import ArchiveMessages +from .archive_media_files import ArchiveMediaFiles -__all__ = ["ArchiveMessages"] +__all__ = ["ArchiveMessages", "ArchiveMediaFiles"] diff --git a/app/repository/archive_media_files_repository.py b/app/repository/archive_media_files_repository.py new file mode 100644 index 0000000..694e8a3 --- /dev/null +++ b/app/repository/archive_media_files_repository.py @@ -0,0 +1,60 @@ +from typing import List, Optional + +from sqlalchemy import select, func +from sqlalchemy.ext.asyncio import AsyncSession + +from app.models.mysql import ArchiveMediaFiles + + +class ArchiveMediaFilesRepository: + def __init__(self, session: AsyncSession): + self.session = session + + async def get_voice_files(self, limit: int = 100, offset: int = 0) -> List[ArchiveMediaFiles]: + """ + 获取语音文件列表 + + Args: + limit: 返回数量限制 + offset: 偏移量 + + Returns: + ArchiveMediaFiles 对象列表 + """ + query = select(ArchiveMediaFiles).where( + ArchiveMediaFiles.file_type == "voice" + ).limit(limit).offset(offset) + + result = await self.session.execute(query) + return list(result.scalars().all()) + + async def count_voice_files(self) -> int: + """ + 统计语音文件总数 + + Returns: + 总数 + """ + query = select(func.count()).select_from(ArchiveMediaFiles).where( + ArchiveMediaFiles.file_type == "voice" + ) + + result = await self.session.execute(query) + return result.scalar() or 0 + + async def get_by_archive_message_id(self, archive_message_id: int) -> Optional[ArchiveMediaFiles]: + """ + 根据 archive_message_id 查询媒体文件 + + Args: + archive_message_id: 归档消息ID + + Returns: + ArchiveMediaFiles 对象或 None + """ + query = select(ArchiveMediaFiles).where( + ArchiveMediaFiles.archive_message_id == archive_message_id + ) + + result = await self.session.execute(query) + return result.scalar_one_or_none() diff --git a/app/repository/archive_messages_repository.py b/app/repository/archive_messages_repository.py index 66cb685..cf3345d 100644 --- a/app/repository/archive_messages_repository.py +++ b/app/repository/archive_messages_repository.py @@ -1,6 +1,6 @@ -from typing import List +from typing import List, Optional -from sqlalchemy import text +from sqlalchemy import text, update from sqlalchemy.ext.asyncio import AsyncSession from app.models.mysql import ArchiveMessages @@ -35,7 +35,7 @@ class ArchiveMessagesRepository: if end_date: conditions.append(" msgtime <= :end_date") params["end_date"] = end_date - base_query += " WHERE 1=1 and roomid='' and msgtype='text' " + base_query += " WHERE 1=1 and roomid='' and msgtype in ('text','voice') " #base_query += " WHERE (from_user='BuShouShiJinBuGaiMing') and roomid='' and msgtype='text' " if conditions: base_query += " AND ".join(conditions)+" " @@ -57,7 +57,7 @@ class ArchiveMessagesRepository: SELECT * FROM archive_messages WHERE DATE(FROM_UNIXTIME(msgtime / 1000)) = :day AND from_user = :from_user - AND to_user = :to_user and msgtype='text' + AND to_user = :to_user and msgtype in ('text','voice') ORDER BY created_at DESC """ result = await self.session.execute(text(sql), {"day": day, "from_user": from_user, "to_user": to_user}) @@ -65,8 +65,32 @@ class ArchiveMessagesRepository: sql = """ SELECT from_user,to_user FROM archive_messages WHERE DATE(FROM_UNIXTIME(msgtime / 1000)) = :day - AND from_user = :from_user and msgtype='text' + AND from_user = :from_user and msgtype in ('text','voice') group by to_user """ result = await self.session.execute(text(sql), {"day": day, "from_user": from_user}) - return [ArchiveMessages(**dict(row)) for row in result.mappings().fetchall()] \ No newline at end of file + return [ArchiveMessages(**dict(row)) for row in result.mappings().fetchall()] + + async def update_content_by_id(self, message_id: int, content: str) -> bool: + """ + 根据ID更新消息的content字段 + + Args: + message_id: 消息ID + content: 新的内容 + + Returns: + 是否更新成功 + """ + try: + stmt = ( + update(ArchiveMessages) + .where(ArchiveMessages.id == message_id) + .values(content=content) + ) + result = await self.session.execute(stmt) + await self.session.commit() + return result.rowcount > 0 + except Exception as e: + await self.session.rollback() + raise e \ No newline at end of file diff --git a/app/repository/redis/__init__.py b/app/repository/redis/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/app/repository/redis/redis_repository.py b/app/repository/redis/redis_repository.py new file mode 100644 index 0000000..0e1fa93 --- /dev/null +++ b/app/repository/redis/redis_repository.py @@ -0,0 +1,283 @@ +import json +import logging +from typing import Any, Optional, List + +from app.client.redis_client_manager import redis_client_manager + +logger = logging.getLogger(__name__) + + +class RedisUtils: + @staticmethod + async def get(key: str) -> Optional[str]: + try: + return await redis_client_manager.client.get(key) + except Exception as e: + logger.error(f"Redis get error: {e}") + return None + + @staticmethod + async def set(key: str, value: str, expire_seconds: Optional[int] = None) -> bool: + try: + if expire_seconds: + await redis_client_manager.client.set(key, value, ex=expire_seconds) + else: + await redis_client_manager.client.set(key, value) + return True + except Exception as e: + logger.error(f"Redis set error: {e}") + return False + + @staticmethod + async def delete(key: str) -> bool: + try: + await redis_client_manager.client.delete(key) + return True + except Exception as e: + logger.error(f"Redis delete error: {e}") + return False + + @staticmethod + async def exists(key: str) -> bool: + try: + return await redis_client_manager.client.exists(key) > 0 + except Exception as e: + logger.error(f"Redis exists error: {e}") + return False + + @staticmethod + async def expire(key: str, seconds: int) -> bool: + try: + return await redis_client_manager.client.expire(key, seconds) + except Exception as e: + logger.error(f"Redis expire error: {e}") + return False + + @staticmethod + async def get_json(key: str) -> Optional[Any]: + try: + value = await redis_client_manager.client.get(key) + if value: + return json.loads(value) + return None + except json.JSONDecodeError: + logger.error(f"Redis get_json error: Invalid JSON") + return None + except Exception as e: + logger.error(f"Redis get_json error: {e}") + return None + + @staticmethod + async def set_json(key: str, value: Any, expire_seconds: Optional[int] = None) -> bool: + try: + json_str = json.dumps(value, ensure_ascii=False) + return await RedisUtils.set(key, json_str, expire_seconds) + except Exception as e: + logger.error(f"Redis set_json error: {e}") + return False + + @staticmethod + async def hget(key: str, field: str) -> Optional[str]: + try: + return await redis_client_manager.client.hget(key, field) + except Exception as e: + logger.error(f"Redis hget error: {e}") + return None + + @staticmethod + async def hset(key: str, field: str, value: str) -> bool: + try: + await redis_client_manager.client.hset(key, field, value) + return True + except Exception as e: + logger.error(f"Redis hset error: {e}") + return False + + @staticmethod + async def hgetall(key: str) -> Optional[dict]: + try: + return await redis_client_manager.client.hgetall(key) + except Exception as e: + logger.error(f"Redis hgetall error: {e}") + return None + + @staticmethod + async def hdel(key: str, field: str) -> bool: + try: + await redis_client_manager.client.hdel(key, field) + return True + except Exception as e: + logger.error(f"Redis hdel error: {e}") + return False + + @staticmethod + async def lpush(key: str, *values: Any) -> int: + try: + return await redis_client_manager.client.lpush(key, *values) + except Exception as e: + logger.error(f"Redis lpush error: {e}") + return 0 + + @staticmethod + async def rpush(key: str, *values: Any) -> int: + try: + return await redis_client_manager.client.rpush(key, *values) + except Exception as e: + logger.error(f"Redis rpush error: {e}") + return 0 + + @staticmethod + async def lpop(key: str) -> Optional[str]: + try: + return await redis_client_manager.client.lpop(key) + except Exception as e: + logger.error(f"Redis lpop error: {e}") + return None + + @staticmethod + async def rpop(key: str) -> Optional[str]: + try: + return await redis_client_manager.client.rpop(key) + except Exception as e: + logger.error(f"Redis rpop error: {e}") + return None + + @staticmethod + async def llen(key: str) -> int: + try: + return await redis_client_manager.client.llen(key) + except Exception as e: + logger.error(f"Redis llen error: {e}") + return 0 + + @staticmethod + async def lrange(key: str, start: int, end: int) -> List[str]: + try: + return await redis_client_manager.client.lrange(key, start, end) + except Exception as e: + logger.error(f"Redis lrange error: {e}") + return [] + + @staticmethod + async def incr(key: str) -> int: + try: + return await redis_client_manager.client.incr(key) + except Exception as e: + logger.error(f"Redis incr error: {e}") + return 0 + + @staticmethod + async def decr(key: str) -> int: + try: + return await redis_client_manager.client.decr(key) + except Exception as e: + logger.error(f"Redis decr error: {e}") + return 0 + + @staticmethod + async def sadd(key: str, *members: Any) -> int: + try: + return await redis_client_manager.client.sadd(key, *members) + except Exception as e: + logger.error(f"Redis sadd error: {e}") + return 0 + + @staticmethod + async def smembers(key: str) -> List[str]: + try: + return await redis_client_manager.client.smembers(key) + except Exception as e: + logger.error(f"Redis smembers error: {e}") + return [] + + @staticmethod + async def srem(key: str, *members: Any) -> int: + try: + return await redis_client_manager.client.srem(key, *members) + except Exception as e: + logger.error(f"Redis srem error: {e}") + return 0 + + @staticmethod + async def zadd(key: str, mapping: dict) -> int: + try: + return await redis_client_manager.client.zadd(key, mapping) + except Exception as e: + logger.error(f"Redis zadd error: {e}") + return 0 + + @staticmethod + async def zrange(key: str, start: int, end: int, withscores: bool = False) -> List[str]: + try: + return await redis_client_manager.client.zrange(key, start, end, withscores=withscores) + except Exception as e: + logger.error(f"Redis zrange error: {e}") + return [] + + @staticmethod + async def keys(pattern: str) -> List[str]: + try: + return await redis_client_manager.client.keys(pattern) + except Exception as e: + logger.error(f"Redis keys error: {e}") + return [] + + @staticmethod + async def flush_db() -> bool: + try: + await redis_client_manager.client.flushdb() + return True + except Exception as e: + logger.error(f"Redis flush_db error: {e}") + return False + + +class RedisSyncUtils: + @staticmethod + def get(key: str) -> Optional[str]: + try: + return redis_client_manager.sync_client.get(key) + except Exception as e: + logger.error(f"Redis sync get error: {e}") + return None + + @staticmethod + def set(key: str, value: str, expire_seconds: Optional[int] = None) -> bool: + try: + if expire_seconds: + redis_client_manager.sync_client.set(key, value, ex=expire_seconds) + else: + redis_client_manager.sync_client.set(key, value) + return True + except Exception as e: + logger.error(f"Redis sync set error: {e}") + return False + + @staticmethod + def delete(key: str) -> bool: + try: + redis_client_manager.sync_client.delete(key) + return True + except Exception as e: + logger.error(f"Redis sync delete error: {e}") + return False + + @staticmethod + def get_json(key: str) -> Optional[Any]: + try: + value = redis_client_manager.sync_client.get(key) + if value: + return json.loads(value) + return None + except Exception as e: + logger.error(f"Redis sync get_json error: {e}") + return None + + @staticmethod + def set_json(key: str, value: Any, expire_seconds: Optional[int] = None) -> bool: + try: + json_str = json.dumps(value, ensure_ascii=False) + return RedisSyncUtils.set(key, json_str, expire_seconds) + except Exception as e: + logger.error(f"Redis sync set_json error: {e}") + return False \ No newline at end of file diff --git a/app/summary/service.py b/app/summary/service.py index b030f29..f67f1a4 100644 --- a/app/summary/service.py +++ b/app/summary/service.py @@ -161,7 +161,11 @@ class SummaryService: [用户] 咨询订单发货状态,订单号 12345,昨日下单。 [销售] 查询后告知正在打包,预计今日发出。 """ - lines = [f"{msg.created_at} {msg.from_role}:{msg.content}" for msg in coro_list] + lines = [ + f"{msg.created_at} {msg.from_role}:{msg.content}" + for msg in coro_list + if not (msg.msgtype == "voice" and "[语音]" not in msg.content) + ] message_str = "\n".join(lines) prompt_template = PromptTemplate(template=prompt, input_variables=["message_str"]) @@ -216,8 +220,8 @@ async def build(day:str): if __name__ == '__main__': # - start = datetime(2026, 6, 23) - end = datetime(2026, 6, 23) + start = datetime(2026, 6, 24) + end = datetime(2026, 6, 24) result = [] temp = start diff --git a/app/uvicorn_main.py b/app/uvicorn_main.py index 3b6b666..3083a91 100644 --- a/app/uvicorn_main.py +++ b/app/uvicorn_main.py @@ -5,12 +5,14 @@ from apscheduler.schedulers.asyncio import AsyncIOScheduler from fastapi import FastAPI from app.chat_qa_query.api.router.query_router import query_qa_router +from app.asr.api.router.asr_router import asr_router from app.core.liftspan import lifespan from app.summary.service import build from app.core.log import logger app = FastAPI(lifespan=lifespan) app.include_router(query_qa_router) +app.include_router(asr_router) # @app.on_event("startup") # async def startup_event(): diff --git a/conf/app_config.yaml b/conf/app_config.yaml index 3fe8535..92177e3 100644 --- a/conf/app_config.yaml +++ b/conf/app_config.yaml @@ -33,3 +33,15 @@ llm: model_name: deepseek-v4-flash api_key: sk-8edc9f25d59643b3b08efea5d81c55a9 base_url: https://api.deepseek.com + +asr: + model_name: qwen3-asr-flash-realtime + api_key: sk-39ceb8d746014b349109e76f893acb18 + base_url: https://dashscope.aliyuncs.com + +redis: + host: 8.159.132.53 + port: 6379 + password: "" + db: 0 + decode_responses: true diff --git a/pyproject.toml b/pyproject.toml index 6316954..ef22f18 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -20,6 +20,7 @@ dependencies = [ "omegaconf>=2.3.0", "pymilvus>=3.0.0", "pyyaml>=6.0.3", + "redis>=8.0.1", "scikit-learn>=1.9.0", "sentence-transformers>=5.6.0", "sqlalchemy>=2.0.50", diff --git a/uv.lock b/uv.lock index 03602a2..19aa977 100644 --- a/uv.lock +++ b/uv.lock @@ -2214,6 +2214,15 @@ wheels = [ { url = "https://pypi.tuna.tsinghua.edu.cn/packages/f1/12/de94a39c2ef588c7e6455cfbe7343d3b2dc9d6b6b2f40c4c6565744c873d/pyyaml-6.0.3-cp314-cp314t-win_arm64.whl", hash = "sha256:ebc55a14a21cb14062aa4162f906cd962b28e2e9ea38f9b4391244cd8de4ae0b", size = 149341, upload-time = "2025-09-25T21:32:56.828Z" }, ] +[[package]] +name = "redis" +version = "8.0.1" +source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } +sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/cc/c3/928b290c2c0ca99ab96eea5b4ff8f30be8112b075301a7d3ba214a3c8c12/redis-8.0.1.tar.gz", hash = "sha256:afc5a7a2f5a084f5b1880dec548dd45be17db7e43c82a30d84f952aefb05cfb0", size = 5114170, upload-time = "2026-06-23T14:52:37.728Z" } +wheels = [ + { url = "https://pypi.tuna.tsinghua.edu.cn/packages/fd/0a/c2345ebf1ebe70840ce3f6c6ee612f8fa749cfbd1b03069c53bf0c62aaad/redis-8.0.1-py3-none-any.whl", hash = "sha256:47daa35a058c23468d6437f17a8c76882cb316b838ef763036af99b96cedd743", size = 502406, upload-time = "2026-06-23T14:52:36.137Z" }, +] + [[package]] name = "regex" version = "2026.5.9" @@ -2469,6 +2478,7 @@ dependencies = [ { name = "omegaconf" }, { name = "pymilvus" }, { name = "pyyaml" }, + { name = "redis" }, { name = "scikit-learn" }, { name = "sentence-transformers" }, { name = "sqlalchemy" }, @@ -2492,6 +2502,7 @@ requires-dist = [ { name = "omegaconf", specifier = ">=2.3.0" }, { name = "pymilvus", specifier = ">=3.0.0" }, { name = "pyyaml", specifier = ">=6.0.3" }, + { name = "redis", specifier = ">=8.0.1" }, { name = "scikit-learn", specifier = ">=1.9.0" }, { name = "sentence-transformers", specifier = ">=5.6.0" }, { name = "sqlalchemy", specifier = ">=2.0.50" },