增加asr能力
This commit is contained in:
parent
927ffb896e
commit
699aea73a6
4
app/asr/__init__.py
Normal file
4
app/asr/__init__.py
Normal file
@ -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"]
|
||||||
112
app/asr/api/asr_dependencies.py
Normal file
112
app/asr/api/asr_dependencies.py
Normal file
@ -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)
|
||||||
15
app/asr/api/router/asr_router.py
Normal file
15
app/asr/api/router/asr_router.py
Normal file
@ -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)
|
||||||
12
app/asr/api/schema/asr_schema.py
Normal file
12
app/asr/api/schema/asr_schema.py
Normal file
@ -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
|
||||||
137
app/asr/asr_client.py
Normal file
137
app/asr/asr_client.py
Normal file
@ -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)
|
||||||
61
app/asr/asr_service.py
Normal file
61
app/asr/asr_service.py
Normal file
@ -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())
|
||||||
129
app/asr/asr_voice_processor.py
Normal file
129
app/asr/asr_voice_processor.py
Normal file
@ -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())
|
||||||
71
app/client/redis_client_manager.py
Normal file
71
app/client/redis_client_manager.py
Normal file
@ -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())
|
||||||
@ -54,6 +54,20 @@ class LLMConfig:
|
|||||||
api_key: str
|
api_key: str
|
||||||
base_url: 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
|
@dataclass
|
||||||
class AppConfig:
|
class AppConfig:
|
||||||
logging: LoggingConfig
|
logging: LoggingConfig
|
||||||
@ -61,6 +75,8 @@ class AppConfig:
|
|||||||
embedding: EmbeddingConfig
|
embedding: EmbeddingConfig
|
||||||
llm: LLMConfig
|
llm: LLMConfig
|
||||||
milvus: MilvusConfig
|
milvus: MilvusConfig
|
||||||
|
asr: ASRConfig
|
||||||
|
redis: RedisConfig
|
||||||
|
|
||||||
|
|
||||||
config_file = Path(__file__).parents[2] / 'conf' / 'app_config.yaml'
|
config_file = Path(__file__).parents[2] / 'conf' / 'app_config.yaml'
|
||||||
|
|||||||
@ -7,6 +7,7 @@ from fastapi import FastAPI
|
|||||||
from app.client.embedding_client_manager import embedding_client
|
from app.client.embedding_client_manager import embedding_client
|
||||||
from app.client.milvus_client_manager import milvus_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.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.core.log import logger
|
||||||
from app.summary.service import build
|
from app.summary.service import build
|
||||||
|
|
||||||
@ -21,24 +22,40 @@ scheduler = AsyncIOScheduler(timezone='Asia/Shanghai')
|
|||||||
async def generate_daily_summary():
|
async def generate_daily_summary():
|
||||||
yesterday = datetime.now() - timedelta(days=1)
|
yesterday = datetime.now() - timedelta(days=1)
|
||||||
date = datetime(yesterday.year, yesterday.month, yesterday.day)
|
date = datetime(yesterday.year, yesterday.month, yesterday.day)
|
||||||
# 按 年-月-日 无补零格式输出
|
|
||||||
date_format = f"{date.year}-{date.month}-{date.day}"
|
date_format = f"{date.year}-{date.month}-{date.day}"
|
||||||
#date_format = "2026-06-23"
|
|
||||||
await build(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} 消息摘要生成完成")
|
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
|
@asynccontextmanager
|
||||||
async def lifespan(app: FastAPI):
|
async def lifespan(app: FastAPI):
|
||||||
db_assistant_mysql_client_manager.init()
|
db_assistant_mysql_client_manager.init()
|
||||||
embedding_client.init()
|
embedding_client.init()
|
||||||
milvus_client.init()
|
milvus_client.init()
|
||||||
|
redis_client_manager.init()
|
||||||
|
|
||||||
scheduler.add_job(
|
scheduler.add_job(
|
||||||
generate_daily_summary,
|
generate_daily_summary,
|
||||||
trigger='cron',
|
trigger='cron',
|
||||||
hour=9,
|
hour=2,
|
||||||
minute=40,
|
minute=0,
|
||||||
second=0,
|
second=0,
|
||||||
id='daily_summary',
|
id='daily_summary',
|
||||||
name='每日消息摘要任务',
|
name='每日消息摘要任务',
|
||||||
@ -50,5 +67,6 @@ async def lifespan(app: FastAPI):
|
|||||||
await db_assistant_mysql_client_manager.close()
|
await db_assistant_mysql_client_manager.close()
|
||||||
await embedding_client.close()
|
await embedding_client.close()
|
||||||
milvus_client.close()
|
milvus_client.close()
|
||||||
|
await redis_client_manager.close()
|
||||||
scheduler.shutdown()
|
scheduler.shutdown()
|
||||||
logger.info("定时任务调度器已关闭")
|
logger.info("定时任务调度器已关闭")
|
||||||
@ -1,3 +1,4 @@
|
|||||||
from .archive_messages import ArchiveMessages
|
from .archive_messages import ArchiveMessages
|
||||||
|
from .archive_media_files import ArchiveMediaFiles
|
||||||
|
|
||||||
__all__ = ["ArchiveMessages"]
|
__all__ = ["ArchiveMessages", "ArchiveMediaFiles"]
|
||||||
|
|||||||
60
app/repository/archive_media_files_repository.py
Normal file
60
app/repository/archive_media_files_repository.py
Normal file
@ -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()
|
||||||
@ -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 sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
from app.models.mysql import ArchiveMessages
|
from app.models.mysql import ArchiveMessages
|
||||||
@ -35,7 +35,7 @@ class ArchiveMessagesRepository:
|
|||||||
if end_date:
|
if end_date:
|
||||||
conditions.append(" msgtime <= :end_date")
|
conditions.append(" msgtime <= :end_date")
|
||||||
params["end_date"] = 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' "
|
#base_query += " WHERE (from_user='BuShouShiJinBuGaiMing') and roomid='' and msgtype='text' "
|
||||||
if conditions:
|
if conditions:
|
||||||
base_query += " AND ".join(conditions)+" "
|
base_query += " AND ".join(conditions)+" "
|
||||||
@ -57,7 +57,7 @@ class ArchiveMessagesRepository:
|
|||||||
SELECT * FROM archive_messages
|
SELECT * FROM archive_messages
|
||||||
WHERE DATE(FROM_UNIXTIME(msgtime / 1000)) = :day
|
WHERE DATE(FROM_UNIXTIME(msgtime / 1000)) = :day
|
||||||
AND from_user = :from_user
|
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
|
ORDER BY created_at DESC
|
||||||
"""
|
"""
|
||||||
result = await self.session.execute(text(sql), {"day": day, "from_user": from_user, "to_user": to_user})
|
result = await self.session.execute(text(sql), {"day": day, "from_user": from_user, "to_user": to_user})
|
||||||
@ -65,8 +65,32 @@ class ArchiveMessagesRepository:
|
|||||||
sql = """
|
sql = """
|
||||||
SELECT from_user,to_user FROM archive_messages
|
SELECT from_user,to_user FROM archive_messages
|
||||||
WHERE DATE(FROM_UNIXTIME(msgtime / 1000)) = :day
|
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
|
group by to_user
|
||||||
"""
|
"""
|
||||||
result = await self.session.execute(text(sql), {"day": day, "from_user": from_user})
|
result = await self.session.execute(text(sql), {"day": day, "from_user": from_user})
|
||||||
return [ArchiveMessages(**dict(row)) for row in result.mappings().fetchall()]
|
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
|
||||||
0
app/repository/redis/__init__.py
Normal file
0
app/repository/redis/__init__.py
Normal file
283
app/repository/redis/redis_repository.py
Normal file
283
app/repository/redis/redis_repository.py
Normal file
@ -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
|
||||||
@ -161,7 +161,11 @@ class SummaryService:
|
|||||||
[用户] 咨询订单发货状态,订单号 12345,昨日下单。
|
[用户] 咨询订单发货状态,订单号 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)
|
message_str = "\n".join(lines)
|
||||||
|
|
||||||
prompt_template = PromptTemplate(template=prompt, input_variables=["message_str"])
|
prompt_template = PromptTemplate(template=prompt, input_variables=["message_str"])
|
||||||
@ -216,8 +220,8 @@ async def build(day:str):
|
|||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
#
|
#
|
||||||
start = datetime(2026, 6, 23)
|
start = datetime(2026, 6, 24)
|
||||||
end = datetime(2026, 6, 23)
|
end = datetime(2026, 6, 24)
|
||||||
|
|
||||||
result = []
|
result = []
|
||||||
temp = start
|
temp = start
|
||||||
|
|||||||
@ -5,12 +5,14 @@ from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
|||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
|
|
||||||
from app.chat_qa_query.api.router.query_router import query_qa_router
|
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.core.liftspan import lifespan
|
||||||
from app.summary.service import build
|
from app.summary.service import build
|
||||||
from app.core.log import logger
|
from app.core.log import logger
|
||||||
|
|
||||||
app = FastAPI(lifespan=lifespan)
|
app = FastAPI(lifespan=lifespan)
|
||||||
app.include_router(query_qa_router)
|
app.include_router(query_qa_router)
|
||||||
|
app.include_router(asr_router)
|
||||||
|
|
||||||
# @app.on_event("startup")
|
# @app.on_event("startup")
|
||||||
# async def startup_event():
|
# async def startup_event():
|
||||||
|
|||||||
@ -33,3 +33,15 @@ llm:
|
|||||||
model_name: deepseek-v4-flash
|
model_name: deepseek-v4-flash
|
||||||
api_key: sk-8edc9f25d59643b3b08efea5d81c55a9
|
api_key: sk-8edc9f25d59643b3b08efea5d81c55a9
|
||||||
base_url: https://api.deepseek.com
|
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
|
||||||
|
|||||||
@ -20,6 +20,7 @@ dependencies = [
|
|||||||
"omegaconf>=2.3.0",
|
"omegaconf>=2.3.0",
|
||||||
"pymilvus>=3.0.0",
|
"pymilvus>=3.0.0",
|
||||||
"pyyaml>=6.0.3",
|
"pyyaml>=6.0.3",
|
||||||
|
"redis>=8.0.1",
|
||||||
"scikit-learn>=1.9.0",
|
"scikit-learn>=1.9.0",
|
||||||
"sentence-transformers>=5.6.0",
|
"sentence-transformers>=5.6.0",
|
||||||
"sqlalchemy>=2.0.50",
|
"sqlalchemy>=2.0.50",
|
||||||
|
|||||||
11
uv.lock
generated
11
uv.lock
generated
@ -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" },
|
{ 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]]
|
[[package]]
|
||||||
name = "regex"
|
name = "regex"
|
||||||
version = "2026.5.9"
|
version = "2026.5.9"
|
||||||
@ -2469,6 +2478,7 @@ dependencies = [
|
|||||||
{ name = "omegaconf" },
|
{ name = "omegaconf" },
|
||||||
{ name = "pymilvus" },
|
{ name = "pymilvus" },
|
||||||
{ name = "pyyaml" },
|
{ name = "pyyaml" },
|
||||||
|
{ name = "redis" },
|
||||||
{ name = "scikit-learn" },
|
{ name = "scikit-learn" },
|
||||||
{ name = "sentence-transformers" },
|
{ name = "sentence-transformers" },
|
||||||
{ name = "sqlalchemy" },
|
{ name = "sqlalchemy" },
|
||||||
@ -2492,6 +2502,7 @@ requires-dist = [
|
|||||||
{ name = "omegaconf", specifier = ">=2.3.0" },
|
{ name = "omegaconf", specifier = ">=2.3.0" },
|
||||||
{ name = "pymilvus", specifier = ">=3.0.0" },
|
{ name = "pymilvus", specifier = ">=3.0.0" },
|
||||||
{ name = "pyyaml", specifier = ">=6.0.3" },
|
{ name = "pyyaml", specifier = ">=6.0.3" },
|
||||||
|
{ name = "redis", specifier = ">=8.0.1" },
|
||||||
{ name = "scikit-learn", specifier = ">=1.9.0" },
|
{ name = "scikit-learn", specifier = ">=1.9.0" },
|
||||||
{ name = "sentence-transformers", specifier = ">=5.6.0" },
|
{ name = "sentence-transformers", specifier = ">=5.6.0" },
|
||||||
{ name = "sqlalchemy", specifier = ">=2.0.50" },
|
{ name = "sqlalchemy", specifier = ">=2.0.50" },
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user