修改设备端对接

This commit is contained in:
HycJack
2026-04-16 21:13:06 +08:00
parent a22fcafea9
commit 602f0f7737
19 changed files with 1385 additions and 60 deletions

Binary file not shown.

Binary file not shown.

Binary file not shown.

View File

@@ -37,6 +37,16 @@ class Settings(BaseSettings):
db_name: str = "talkingq"
db_echo: bool = False # 是否打印SQL语句
# TalkingQ设备MQTT命令服务配置
talkingq_mqtt_broker: str = "broker.emqx.io"
talkingq_mqtt_port: int = 1883
talkingq_mqtt_username: str = ""
talkingq_mqtt_password: str = ""
talkingq_mqtt_device_prefix: str = "TalkingQ"
talkingq_mqtt_qos: int = 1
talkingq_mqtt_keepalive: int = 60
talkingq_mqtt_nfc_notice_interval: int = 600
admin_api_key: str # 用于设备注册的管理员API密钥
client_api_key: str # 用于微信小程序客户端验证的API密钥

View File

@@ -126,3 +126,19 @@ class SystemConfig(Base):
__table_args__ = (
{'mysql_charset': 'utf8mb4', 'mysql_collate': 'utf8mb4_unicode_ci'}
)
class Card(Base):
__tablename__ = "cards"
card_id = Column(Integer, primary_key=True, autoincrement=True)
card_uuid = Column(String(64), unique=True, nullable=False)
device_id = Column(String(64), index=True)
card_name = Column(String(64))
status = Column(Integer, server_default="0")
total_swaps = Column(Integer, server_default="0")
created_at = Column(DateTime, server_default=func.now())
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
__table_args__ = (
{'mysql_charset': 'utf8mb4', 'mysql_collate': 'utf8mb4_unicode_ci'}
)

View File

@@ -0,0 +1,36 @@
import os
import uuid
from config import settings
from utils.logger import session_logger
from utils.audio_denoiser import reduce_background_noise
async def save_audio_file(audio_data: bytes, device_id: str) -> str:
"""
保存音频数据到 assets/audio 目录
Args:
audio_data: 音频二进制数据
device_id: 设备ID
Returns:
音频文件的相对路径
"""
try:
audio_dir = os.path.join(settings.assets_dir, "audio")
os.makedirs(audio_dir, exist_ok=True)
filename = f"{device_id}_{uuid.uuid4().hex[:8]}.mp3"
filepath = os.path.join(audio_dir, filename)
with open(filepath, 'wb') as f:
f.write(audio_data)
# relative_path = f"assets/audio/{filename}"
session_logger.info(device_id, "audio", f"音频文件已保存: {filepath}")
# reduce_background_noise(filepath, relative_path,noise_path='assets/audio/noise_sample.wav',normalize_volume=True)
return filepath
except Exception as e:
session_logger.error(device_id, "audio", f"保存音频文件时出错: {e}", exc_info=True)
raise

View File

@@ -2,6 +2,7 @@ from fastapi import WebSocket, WebSocketDisconnect
from utils.logger import session_logger
from services.connection_manager import connection_manager
from services.task_manager import task_manager
from services.offline_audio_cache import offline_audio_cache
from handlers.websocket_auth_handler import authenticate_websocket
from handlers.websocket_message_handler import handle_websocket_messages
from handlers.session_cleanup_handler import cleanup_device_sessions
@@ -29,6 +30,20 @@ async def websocket_endpoint(websocket: WebSocket):
await websocket.send_text('{"status": "error", "message": "Not authenticated"}')
return
# 检查设备是否有离线音频URL需要发送
has_pending = await offline_audio_cache.has_pending_audio(device_id)
if has_pending:
audio_urls = await offline_audio_cache.get_audio_urls(device_id)
for audio_url in audio_urls:
try:
await websocket.send_text(f"SOUND_URL:{audio_url}")
session_logger.info(device_id, "offline", f"发送离线音频URL: {audio_url}")
except Exception as e:
session_logger.error(device_id, "offline", f"发送离线音频URL失败: {e}")
# 清空已发送的离线音频URL
await offline_audio_cache.clear_audio_urls(device_id)
session_logger.info(device_id, "offline", f"已清空设备的离线音频URL缓存{len(audio_urls)}")
await handle_websocket_messages(websocket, device_id)
except WebSocketDisconnect:
@@ -43,7 +58,7 @@ async def websocket_endpoint(websocket: WebSocket):
)
finally:
if device_id:
await connection_manager.remove_connection(device_id)
# await connection_manager.remove_connection(device_id)
await cleanup_device_sessions(device_id)
# 清理设备相关的所有异步任务
await task_manager.cancel_device_tasks(device_id)

View File

@@ -4,12 +4,19 @@ import asyncio
from fastapi import WebSocket
from handlers.audio_packet_parser import parse_packet
from handlers.audio_session_handler import handle_websocket_data
from handlers.audio_file_handler import save_audio_file
from services.audio_session import audio_session_manager
from services.interrupt_handler import interrupt_handler
from services.task_manager import task_manager
from services.connection_manager import connection_manager
from services.device_target_cache import device_target_cache
from services.offline_audio_cache import offline_audio_cache
from services.target_audio_cache import target_audio_cache
from services.card_service import card_service
from utils.logger import session_logger
from handlers.prompt_sound_handler import handle_prompt_sound_request
from handlers.session_cleanup_handler import handle_old_session_cleanup
from config import settings
async def handle_websocket_messages(websocket: WebSocket, device_id: str):
"""
@@ -52,7 +59,36 @@ async def handle_websocket_messages(websocket: WebSocket, device_id: str):
async def handle_text_message(websocket: WebSocket, device_id: str, text_data: str):
"""处理文本消息"""
if text_data.startswith("REQUEST_PROMPT_SOUND:"):
if text_data.startswith("REGISTER_TARGET_DEVICE:"):
card_uuid = text_data.split(":", 1)[1].strip()
# 检查卡片是否存在
existing_card = await card_service.get_card_by_uuid(card_uuid)
if existing_card:
# 卡片已存在使用卡片绑定的设备ID作为目标设备ID
target_device_id = existing_card.device_id
session_logger.info(device_id, "card", f"卡片已存在绑定的设备ID: {target_device_id}")
else:
# 卡片不存在,创建新卡片并绑定到当前设备
new_card = await card_service.activate_card(card_uuid, device_id)
session_logger.info(device_id, "card", f"新卡片{card_uuid}已创建并激活,绑定到设备: {device_id}")
return
# 设置目标设备
await device_target_cache.set_target(device_id, target_device_id)
# 检查目标设备是否在线
target_websocket = await connection_manager.get_connection(target_device_id)
if target_websocket and target_websocket.client_state.name == "CONNECTED":
sound_url = f"http://{settings.server_host}:{settings.server_port}/assets/audio/welcome.mp3"
await websocket.send_text(f"TARGET_DEVICE_REGISTERED_URL:{sound_url}")
session_logger.info(device_id, "target", f"成功注册目标设备: {target_device_id}")
else:
sound_url = f"http://{settings.server_host}:{settings.server_port}/assets/audio/offline.mp3"
await websocket.send_text(f"TARGET_DEVICE_REGISTERED_URL:{sound_url}")
session_logger.warning(device_id, "target", f"目标设备 {target_device_id} 不在线")
elif text_data.startswith("REQUEST_PROMPT_SOUND:"):
prompt_type = text_data.split(":", 1)[1]
if device_id:
await handle_prompt_sound_request(device_id, prompt_type)
@@ -107,16 +143,36 @@ async def handle_binary_message(websocket: WebSocket, device_id: str, binary_dat
session_id = session_key[1]
if packet_type == 1: # 开始包
target_device_id = await device_target_cache.get_target(device_id)
if target_device_id:
# 清除音频缓存
await target_audio_cache.clear_audio_data(target_device_id)
return None, None
current_active_session = await handle_start_packet(
device_id, session_id, session_key, session, current_active_session
)
elif packet_type == 3: # 中断包
target_device_id = await device_target_cache.get_target(device_id)
if target_device_id:
return None, None
current_active_session = await handle_interrupt_packet(
device_id, session_id, session_key, session, current_active_session
)
elif packet_type == 2: # 结束包
target_device_id = await device_target_cache.get_target(device_id)
if target_device_id:
# 处理缓存的音频数据
session_logger.info(device_id, session_id, f"收到结束包,开始处理缓存音频数据")
await process_cached_audio(device_id, target_device_id)
# 移除目标设备关联
await device_target_cache.remove_target(device_id)
return None, None
if current_active_session == session_key:
current_active_session = None
elif packet_type == 4: # 发送给目标设备的音频包
await handle_target_audio_packet(device_id, audio_data)
return None, None
elif packet_type == 0 and first_audio_received_time is None:
first_audio_received_time = time.perf_counter()
session.start_time = first_audio_received_time
@@ -178,3 +234,53 @@ async def handle_interrupt_packet(device_id, session_id, session_key, session, c
task_type="interrupt"
)
return current_active_session
async def handle_target_audio_packet(device_id: str, audio_data: bytes):
"""处理发送给目标设备的音频包"""
try:
target_device_id = await device_target_cache.get_target(device_id)
if not target_device_id:
session_logger.warning(device_id, "target", "未设置目标设备,无法发送音频")
return
# 缓存音频数据
await target_audio_cache.add_audio_data(target_device_id, audio_data)
session_logger.info(device_id, "target", f"已缓存音频数据到目标设备 {target_device_id}")
except Exception as e:
session_logger.error(device_id, "target", f"处理目标音频包时出错: {e}", exc_info=True)
async def process_cached_audio(device_id: str, target_device_id: str):
"""处理缓存的音频数据并发送音频URL"""
try:
# 获取缓存的音频数据
cached_audio = await target_audio_cache.get_audio_data(target_device_id)
if not cached_audio:
session_logger.info(device_id, "target", f"目标设备 {target_device_id} 没有缓存的音频数据")
return
# 保存音频文件
audio_path = await save_audio_file(cached_audio, device_id)
audio_url = f"http://{settings.server_host}:{settings.server_port}/{audio_path}"
# 发送URL给目标设备
target_websocket = await connection_manager.get_connection(target_device_id)
if target_websocket and target_websocket.client_state.name == "CONNECTED":
await target_websocket.send_text("TTS_START")
session_logger.info(device_id, "target", "已发送 TTS_START 给客户端")
await target_websocket.send_text(f"NFC_SOUND_URL:{audio_url}")
session_logger.info(device_id, "target", f"已发送音频URL给目标设备 {target_device_id}: {audio_url}")
await target_websocket.send_text("TTS_END")
session_logger.info(device_id, "target", "已发送 TTS_END 给客户端")
else:
# 目标设备不在线,保存到离线缓存
await offline_audio_cache.add_audio_url(target_device_id, audio_url)
session_logger.warning(device_id, "target", f"目标设备 {target_device_id} 不在线保存音频URL到离线缓存")
except Exception as e:
session_logger.error(device_id, "target", f"处理缓存音频时出错: {e}", exc_info=True)
finally:
# 清除缓存
await target_audio_cache.clear_audio_data(target_device_id)

View File

@@ -1,4 +1,5 @@
import asyncio
import threading
from dashscope.audio.asr import TranslationRecognizerChat
from dashscope.audio.asr import TranscriptionResult
from dashscope.audio.asr import TranslationRecognizerCallback
@@ -10,15 +11,18 @@ from utils.logger import session_logger
class GummyCallback(TranslationRecognizerCallback):
def __init__(self, device_id, session_id):
def __init__(self, device_id, session_id, asr_instance=None):
self.device_id = device_id
self.session_id = session_id
self.asr_instance = asr_instance
self.transcript = ""
self.transcript_event = asyncio.Event()
self.received_result = False
self.error_message = None
self.is_closed = False
self.session_end = False
self.reconnect_event = threading.Event()
self.is_rate_limit_error = False
def on_open(self) -> None:
session_logger.info(self.device_id, self.session_id, "Gummy ASR连接建立成功")
@@ -50,9 +54,22 @@ class GummyCallback(TranslationRecognizerCallback):
def on_error(self, result) -> None:
self.error_message = str(result)
session_logger.error(
self.device_id, self.session_id, f"Gummy ASR错误: {self.error_message}"
)
error_str = str(result).lower()
self.is_rate_limit_error = "throttling" in error_str or "rate limit" in error_str
if self.is_rate_limit_error:
session_logger.warning(
self.device_id,
self.session_id,
f"Gummy ASR限流错误触发重连信号: {self.error_message}"
)
self.reconnect_event.set()
if self.asr_instance:
self.asr_instance.need_reconnect = True
else:
session_logger.error(
self.device_id, self.session_id, f"Gummy ASR错误: {self.error_message}"
)
self.transcript_event.set()
def on_close(self) -> None:
@@ -74,8 +91,11 @@ class AliyunASR(ASR, BaseService):
self.audio_buffer = bytearray()
self.translator = None
self.callback = None
self.buffer_size = 3200 # 约100ms的音频按16000采样率计算
self.max_end_silence = 2000 # 设为2000ms允许更长的停顿
self.buffer_size = 3200
self.max_end_silence = 2000
self.max_retries = 3
self.retry_delay = 1.0
self.need_reconnect = False
async def connect(self):
try:
@@ -86,9 +106,8 @@ class AliyunASR(ASR, BaseService):
return
self.transcript = ""
self.audio_buffer = bytearray()
self.callback = GummyCallback(self.device_id, self.session_id)
self.callback = GummyCallback(self.device_id, self.session_id, self)
self.translator = TranslationRecognizerChat(
model="gummy-chat-v1",
@@ -106,20 +125,79 @@ class AliyunASR(ASR, BaseService):
self.translator.start()
self.callback.transcript_event.clear()
self.connected = True
self.need_reconnect = False
except Exception as e:
session_logger.error(
self.device_id, self.session_id, f"建立Gummy ASR连接失败: {str(e)}"
)
self.connected = False
if self.translator:
try:
self.translator.stop()
except:
pass
self.translator = None
await self._cleanup_connection()
async def _cleanup_connection(self):
if self.translator:
try:
self.translator.stop()
except:
pass
self.translator = None
self.callback = None
self.connected = False
async def _handle_reconnect(self):
if not self.need_reconnect:
return True
session_logger.warning(
self.device_id,
self.session_id,
"检测到限流错误,准备重连..."
)
await self._cleanup_connection()
current_delay = self.retry_delay
for attempt in range(self.max_retries):
session_logger.info(
self.device_id,
self.session_id,
f"重连尝试 {attempt + 1}/{self.max_retries},等待{current_delay}"
)
await asyncio.sleep(current_delay)
try:
await self.connect()
if self.connected:
session_logger.info(
self.device_id,
self.session_id,
"重连成功"
)
return True
except Exception as e:
session_logger.error(
self.device_id,
self.session_id,
f"重连失败: {str(e)}"
)
current_delay *= 2
self.retry_delay = current_delay
session_logger.error(
self.device_id,
self.session_id,
"重连失败,已达到最大重试次数"
)
self.need_reconnect = False
return False
async def send_audio(self, audio_data: bytes):
if self.need_reconnect:
reconnected = await self._handle_reconnect()
if not reconnected:
session_logger.error(self.device_id, self.session_id, "Gummy ASR重连失败无法发送音频")
return
if not self.connected or not self.translator:
session_logger.error(self.device_id, self.session_id, "Gummy ASR连接未建立")
return
@@ -128,6 +206,19 @@ class AliyunASR(ASR, BaseService):
self.audio_buffer.extend(audio_data)
while len(self.audio_buffer) >= self.buffer_size:
if self.need_reconnect:
session_logger.info(
self.device_id,
self.session_id,
"检测到限流错误,暂停发送音频,准备重连"
)
reconnected = await self._handle_reconnect()
if not reconnected:
return
if not self.connected or not self.translator:
return
continue
chunk = bytes(self.audio_buffer[: self.buffer_size])
self.audio_buffer = self.audio_buffer[self.buffer_size :]

View File

@@ -24,9 +24,19 @@ class BaseLLM(LLM, BaseService):
async def connect(self):
try:
if self.session:
await self.session.close()
self.session = aiohttp.ClientSession()
if self.session and not self.session.closed:
return
connector = aiohttp.TCPConnector(
limit=20,
limit_per_host=5,
force_close=True,
enable_cleanup_closed=True
)
timeout = aiohttp.ClientTimeout(total=120, connect=30)
self.session = aiohttp.ClientSession(
connector=connector,
timeout=timeout
)
self.connected = True
except Exception as e:
session_logger.error(
@@ -97,8 +107,10 @@ class BaseLLM(LLM, BaseService):
}
try:
if not self.session or self.session.closed:
self.session = aiohttp.ClientSession()
self.connected = True
await self.connect()
if not self.connected:
yield None
return
response = await self.session.post(
f"{self.base_url}/chat/completions", headers=headers, json=data
)

View File

@@ -25,9 +25,24 @@ class MiniMaxTTS(TTS, BaseService):
self.voice_id = selected_role["minimax_voice_id"]
self.model = "speech-02-turbo"
self.selected_role = selected_role
self._client_session: Optional[aiohttp.ClientSession] = None
async def connect(self):
"""连接到MiniMax服务"""
if self._client_session is None or self._client_session.closed:
connector = aiohttp.TCPConnector(
limit=20,
limit_per_host=5,
force_close=True,
enable_cleanup_closed=True
)
timeout = aiohttp.ClientTimeout(total=120, connect=30)
self._client_session = aiohttp.ClientSession(
connector=connector,
timeout=timeout,
max_field_size=1024*1024,
max_line_size=10*1024*1024
)
self.connected = True
session_logger.info(
self.device_id,
@@ -38,6 +53,8 @@ class MiniMaxTTS(TTS, BaseService):
async def close(self):
"""关闭MiniMax服务连接"""
self.connected = False
if self._client_session and not self._client_session.closed:
await self._client_session.close()
async def tts(
self,
@@ -138,48 +155,66 @@ class MiniMaxTTS(TTS, BaseService):
audio_buffer = bytearray()
async with aiohttp.ClientSession() as client_session:
async with client_session.post(url, json=payload, headers=headers) as response:
if response.status != 200:
error_text = await response.text()
session_logger.error(
self.device_id,
self.session_id,
f"MiniMax TTS请求失败: {response.status}, {error_text}"
)
return None
session_logger.info(
if self._client_session is None or self._client_session.closed:
await self.connect()
async with self._client_session.post(url, json=payload, headers=headers) as response:
if response.status != 200:
error_text = await response.text()
session_logger.error(
self.device_id,
self.session_id,
"开始接收MiniMax TTS流式响应"
f"MiniMax TTS请求失败: {response.status}, {error_text}"
)
return None
session_logger.info(
self.device_id,
self.session_id,
"开始接收MiniMax TTS流式响应"
)
buffer = b""
chunk_size = 8192
async for chunk in response.content.iter_chunked(chunk_size):
buffer += chunk
async for line in response.content:
if line.startswith(b'data:'):
try:
data_json = json.loads(line[5:])
if "data" in data_json and "audio" in data_json["data"]:
status = data_json["data"].get("status", 1)
if status == 1: # 只处理status=1(合成中)的音频数据忽略status=2(合成结束)的汇总数据
audio_hex = data_json["data"]["audio"]
audio_binary = binascii.unhexlify(audio_hex)
audio_buffer.extend(audio_binary)
if status == 2: # 合成结束
session_logger.info(
self.device_id,
self.session_id,
"MiniMax TTS流式合成完成"
)
except Exception as e:
session_logger.error(
self.device_id,
self.session_id,
f"处理MiniMax TTS流式响应出错: {str(e)}",
exc_info=True
)
while b'\n' in buffer:
line, buffer = buffer.split(b'\n', 1)
line = line.strip()
if not line or not line.startswith(b'data:'):
continue
try:
data_json = json.loads(line[5:])
if "data" in data_json and "audio" in data_json["data"]:
status = data_json["data"].get("status", 1)
if status == 1:
audio_hex = data_json["data"]["audio"]
audio_binary = binascii.unhexlify(audio_hex)
audio_buffer.extend(audio_binary)
if status == 2:
session_logger.info(
self.device_id,
self.session_id,
"MiniMax TTS流式合成完成"
)
except json.JSONDecodeError as e:
session_logger.warning(
self.device_id,
self.session_id,
f"JSON解析失败: {str(e)}"
)
except Exception as e:
session_logger.error(
self.device_id,
self.session_id,
f"处理MiniMax TTS流式响应出错: {str(e)}",
exc_info=True
)
async with aiofiles.open(output_file, 'wb') as f:
await f.write(audio_buffer)

View File

@@ -11,3 +11,11 @@ aiomysql==0.2.0
sqlalchemy[asyncio]==2.0.40
loguru==0.7.3
noisereduce==3.0.2
soundfile==0.12.1
numpy==1.26.4
scipy==1.11.4
librosa==0.10.1
setuptools==69.5.1
paho-mqtt>=2.0.0

View File

@@ -0,0 +1,190 @@
import asyncio
from typing import Optional, Dict
import time
from sqlalchemy import select, update, insert
from sqlalchemy.ext.asyncio import AsyncSession
from utils.logger import session_logger
from pydantic import BaseModel
from database.models import Card as DBCard
from services.database_service_base import DatabaseServiceBase
class Card(BaseModel):
card_id: Optional[int] = None
card_uuid: str
device_id: Optional[str] = None
card_name: Optional[str] = None
status: int = 0
total_swaps: int = 0
created_at: Optional[float] = None
updated_at: Optional[float] = None
@classmethod
def from_db_model(cls, db_model: DBCard):
"""从数据库模型创建卡片对象"""
return cls(
card_id=db_model.card_id,
card_uuid=db_model.card_uuid,
device_id=db_model.device_id,
card_name=db_model.card_name,
status=db_model.status,
total_swaps=db_model.total_swaps,
created_at=db_model.created_at.timestamp() if db_model.created_at else None,
updated_at=db_model.updated_at.timestamp() if db_model.updated_at else None
)
class CardService(DatabaseServiceBase):
def __init__(self):
super().__init__(service_name="card")
self.cards: Dict[str, Card] = {}
self.lock = asyncio.Lock()
async def _load_card_from_db(self, card_uuid: str, async_session: AsyncSession) -> Optional[Card]:
"""从数据库加载卡片信息"""
try:
query = select(DBCard).where(DBCard.card_uuid == card_uuid)
result = await async_session.execute(query)
db_card = result.scalar_one_or_none()
if db_card:
card = Card.from_db_model(db_card)
async with self.lock:
self.cards[card_uuid] = card
return card
return None
except Exception as e:
session_logger.error("card", "service", f"从数据库加载卡片失败: {str(e)}")
return None
async def _save_card_to_db(self, card: Card, async_session: AsyncSession):
"""保存卡片信息到数据库"""
try:
query = select(DBCard).where(DBCard.card_uuid == card.card_uuid)
result = await async_session.execute(query)
existing_card = result.scalar_one_or_none()
if existing_card:
stmt = update(DBCard).where(
DBCard.card_uuid == card.card_uuid
).values(
device_id=card.device_id,
card_name=card.card_name,
status=card.status,
total_swaps=card.total_swaps
)
else:
stmt = insert(DBCard).values(
card_uuid=card.card_uuid,
device_id=card.device_id,
card_name=card.card_name,
status=card.status,
total_swaps=card.total_swaps
)
await async_session.execute(stmt)
await async_session.commit()
session_logger.info("card", "service", f"卡片已保存到数据库: {card.card_uuid}")
except Exception as e:
await async_session.rollback()
session_logger.error("card", "service", f"保存卡片到数据库失败: {str(e)}")
raise
async def get_card_by_uuid(self, card_uuid: str, force_refresh: bool = False) -> Optional[Card]:
"""根据UUID获取卡片信息"""
await self._init_database()
card = None
if not force_refresh:
async with self.lock:
card = self.cards.get(card_uuid)
if not card:
db_session = await self.db_manager.get_session()
try:
card = await self._load_card_from_db(card_uuid, db_session)
finally:
await db_session.close()
return card
async def get_card_by_device_id(self, device_id: str) -> list[Card]:
"""根据设备ID获取卡片列表"""
await self._init_database()
db_session = await self.db_manager.get_session()
try:
query = select(DBCard).where(DBCard.device_id == device_id)
result = await db_session.execute(query)
db_cards = result.scalars().all()
cards = []
for db_card in db_cards:
card = Card.from_db_model(db_card)
async with self.lock:
self.cards[card.card_uuid] = card
cards.append(card)
return cards
except Exception as e:
session_logger.error(device_id, "card", f"根据设备ID获取卡片失败: {str(e)}")
return []
finally:
await db_session.close()
async def activate_card(self, card_uuid: str, device_id: str, card_name: Optional[str] = None) -> Card:
"""激活卡片并绑定到设备"""
await self._init_database()
db_session = await self.db_manager.get_session()
try:
# 检查卡片是否已存在
existing_card = await self.get_card_by_uuid(card_uuid)
if existing_card:
# 更新现有卡片
existing_card.device_id = device_id
existing_card.card_name = card_name
existing_card.status = 1 # 激活状态
await self._save_card_to_db(existing_card, db_session)
session_logger.info(device_id, "card", f"卡片已激活并绑定到设备: {card_uuid}")
return existing_card
else:
# 创建新卡片
new_card = Card(
card_uuid=card_uuid,
device_id=device_id,
card_name=card_name,
status=1, # 激活状态
total_swaps=0
)
await self._save_card_to_db(new_card, db_session)
session_logger.info(device_id, "card", f"新卡片已创建并激活: {card_uuid}")
return new_card
finally:
await db_session.close()
async def increment_swap_count(self, card_uuid: str) -> Optional[Card]:
"""增加卡片交换次数"""
await self._init_database()
card = await self.get_card_by_uuid(card_uuid)
if card:
card.total_swaps += 1
db_session = await self.db_manager.get_session()
try:
await self._save_card_to_db(card, db_session)
session_logger.info("card", "service", f"卡片交换次数已增加: {card_uuid}, 总次数: {card.total_swaps}")
return card
finally:
await db_session.close()
return None
async def check_card_ownership(self, card_uuid: str, device_id: str) -> bool:
"""检查卡片是否属于指定设备"""
card = await self.get_card_by_uuid(card_uuid)
if card and card.device_id == device_id:
return True
return False
card_service = CardService()

View File

@@ -0,0 +1,29 @@
import asyncio
from typing import Dict, Optional
class DeviceTargetCache:
def __init__(self):
self.targets: Dict[str, str] = {}
self.lock = asyncio.Lock()
async def get_target(self, device_id: str) -> Optional[str]:
async with self.lock:
return self.targets.get(device_id)
async def set_target(self, device_id: str, target_device_id: str):
async with self.lock:
self.targets[device_id] = target_device_id
async def remove_target(self, device_id: str):
async with self.lock:
if device_id in self.targets:
del self.targets[device_id]
async def get_all_targets(self) -> list:
async with self.lock:
return list(self.targets.items())
# 单例实例
device_target_cache = DeviceTargetCache()

View File

@@ -0,0 +1,31 @@
import asyncio
from typing import Dict, List
class OfflineAudioCache:
def __init__(self):
self.offline_audio: Dict[str, List[str]] = {}
self.lock = asyncio.Lock()
async def add_audio_url(self, device_id: str, audio_url: str):
async with self.lock:
if device_id not in self.offline_audio:
self.offline_audio[device_id] = []
self.offline_audio[device_id].append(audio_url)
async def get_audio_urls(self, device_id: str) -> List[str]:
async with self.lock:
return self.offline_audio.get(device_id, [])
async def clear_audio_urls(self, device_id: str):
async with self.lock:
if device_id in self.offline_audio:
del self.offline_audio[device_id]
async def has_pending_audio(self, device_id: str) -> bool:
async with self.lock:
return device_id in self.offline_audio and len(self.offline_audio[device_id]) > 0
# 单例实例
offline_audio_cache = OfflineAudioCache()

View File

@@ -0,0 +1,35 @@
import asyncio
from typing import Dict, List, Optional
class TargetAudioCache:
def __init__(self):
self.audio_data_cache: Dict[str, List[bytes]] = {}
self.lock = asyncio.Lock()
async def add_audio_data(self, device_id: str, audio_data: bytes):
async with self.lock:
if device_id not in self.audio_data_cache:
self.audio_data_cache[device_id] = []
self.audio_data_cache[device_id].append(audio_data)
async def get_audio_data(self, device_id: str) -> Optional[bytes]:
async with self.lock:
if device_id not in self.audio_data_cache:
return None
# 合并所有音频数据
combined_audio = b''.join(self.audio_data_cache[device_id])
return combined_audio
async def clear_audio_data(self, device_id: str):
async with self.lock:
if device_id in self.audio_data_cache:
del self.audio_data_cache[device_id]
async def has_cached_audio(self, device_id: str) -> bool:
async with self.lock:
return device_id in self.audio_data_cache and len(self.audio_data_cache[device_id]) > 0
# 单例实例
target_audio_cache = TargetAudioCache()

View File

@@ -0,0 +1,538 @@
import asyncio
from test.audio_processor import AudioProcessor
import math
import asyncio
import struct
import websockets
from dashscope.audio.asr import TranslationRecognizerChat
from dashscope.audio.asr import TranscriptionResult
from dashscope.audio.asr import TranslationRecognizerCallback
# from config import settings
from providers.asr import ASR
# from providers.impl.base_service import BaseService
from providers.impl.interrupt_handler import interrupt_handler
from utils.logger import session_logger
class GummyCallback(TranslationRecognizerCallback):
def __init__(self, device_id, session_id):
self.device_id = device_id
self.session_id = session_id
self.transcript = ""
self.transcript_event = asyncio.Event()
self.received_result = False
self.error_message = None
self.is_closed = False
self.session_end = False
def on_open(self) -> None:
session_logger.info(self.device_id, self.session_id, "Gummy ASR连接建立成功")
def on_event(
self,
request_id,
transcription_result: TranscriptionResult,
translation_result,
usage,
) -> None:
if transcription_result:
if transcription_result.is_sentence_end:
self.transcript = transcription_result.text
self.received_result = True
session_logger.info(
self.device_id,
self.session_id,
f"Gummy ASR完整句子: {self.transcript}",
)
self.transcript_event.set()
else:
self.transcript = transcription_result.text
def on_complete(self) -> None:
session_logger.info(self.device_id, self.session_id, "Gummy ASR识别完成")
self.session_end = True
self.transcript_event.set()
def on_error(self, result) -> None:
self.error_message = str(result)
session_logger.error(
self.device_id, self.session_id, f"Gummy ASR错误: {self.error_message}"
)
self.transcript_event.set()
def on_close(self) -> None:
session_logger.info(self.device_id, self.session_id, "Gummy ASR连接已关闭")
self.is_closed = True
self.transcript_event.set()
class AliyunASR(ASR):
def __init__(self, selected_role=None):
# BaseService.__init__(self, "asr", selected_role)
self.api_key = 'sk-7a50eca6856d4afb968ac3bf512f6d1b'
self.vocabulary_id = 'vocab-talkingq-8c1ed0b3b76342819a1d39988ff22a17'
self.connected = False
self.transcript = ""
self.device_id = "unknown"
self.session_id = "unknown"
self.sample_rate = 16000
self.audio_buffer = bytearray()
self.translator = None
self.callback = None
self.buffer_size = 3200 # 约100ms的音频按16000采样率计算
self.max_end_silence = 2000 # 设为2000ms允许更长的停顿
async def connect(self):
try:
if self.connected and self.translator:
session_logger.info(
self.device_id, self.session_id, "Gummy ASR连接已存在复用当前连接"
)
return
self.transcript = ""
self.audio_buffer = bytearray()
self.callback = GummyCallback(self.device_id, self.session_id)
self.translator = TranslationRecognizerChat(
model="gummy-chat-v1",
format="mp3",
sample_rate=self.sample_rate,
callback=self.callback,
transcription_enabled=True,
source_language="auto",
semantic_punctuation_enabled=False,
max_end_silence=self.max_end_silence,
api_key=self.api_key,
vocabulary_id=self.vocabulary_id,
)
self.translator.start()
self.callback.transcript_event.clear()
self.connected = True
except Exception as e:
session_logger.error(
self.device_id, self.session_id, f"建立Gummy ASR连接失败: {str(e)}"
)
self.connected = False
if self.translator:
try:
self.translator.stop()
except:
pass
self.translator = None
async def send_audio(self, audio_data: bytes):
if not self.connected or not self.translator:
session_logger.error(self.device_id, self.session_id, "Gummy ASR连接未建立")
return
try:
self.audio_buffer.extend(audio_data)
while len(self.audio_buffer) >= self.buffer_size:
chunk = bytes(self.audio_buffer[: self.buffer_size])
self.audio_buffer = self.audio_buffer[self.buffer_size :]
if self.device_id and self.session_id:
session_key = (self.device_id, self.session_id)
if interrupt_handler.is_interrupted(session_key):
session_logger.info(
self.device_id,
self.session_id,
"检测到中断,停止发送音频数据",
)
return
try:
self.translator.send_audio_frame(chunk)
except Exception as e:
session_logger.error(
self.device_id,
self.session_id,
f"发送音频帧到Gummy ASR出错: {str(e)}",
)
except Exception as e:
session_logger.error(
self.device_id, self.session_id, f"发送音频到Gummy ASR出错: {str(e)}"
)
async def send_end(self):
if not self.connected or not self.translator:
session_logger.error(
self.device_id, self.session_id, "Gummy ASR连接未建立无法发送结束信号"
)
return
try:
if len(self.audio_buffer) > 0:
chunk = bytes(self.audio_buffer)
self.audio_buffer = bytearray()
self.translator.send_audio_frame(chunk)
session_logger.info(
self.device_id,
self.session_id,
"客户端发送结束包等待Gummy ASR处理完毕",
)
loop = asyncio.get_running_loop()
await loop.run_in_executor(None, lambda: self._safe_stop_translator())
except Exception as e:
session_logger.error(
self.device_id,
self.session_id,
f"发送结束信号到Gummy ASR出错: {str(e)}",
)
async def receive_results(self):
if not self.connected:
session_logger.error(
self.device_id, self.session_id, "Gummy ASR连接未建立无法接收结果"
)
return
try:
try:
await asyncio.wait_for(
self.callback.transcript_event.wait(), timeout=0.5
)
except asyncio.TimeoutError:
session_logger.warning(
self.device_id,
self.session_id,
"接收Gummy ASR结果超时使用最后的中间结果",
)
if self.callback.received_result or self.callback.transcript:
self.transcript = self.callback.transcript
except Exception as e:
session_logger.error(
self.device_id, self.session_id, f"接收Gummy ASR结果出错: {str(e)}"
)
return self.transcript
async def close(self):
if self.connected:
try:
if self.translator is not None:
try:
self.translator.stop()
except Exception as e:
if "has stopped" not in str(e):
session_logger.warning(
self.device_id,
self.session_id,
f"关闭Gummy translator出错: {str(e)}"
)
self.translator = None
self.connected = False
session_logger.info(
self.device_id, self.session_id, "Gummy ASR连接已关闭"
)
except Exception as e:
session_logger.error(
self.device_id, self.session_id, f"关闭Gummy ASR连接出错: {str(e)}"
)
def _safe_stop_translator(self):
"""安全停止translator忽略已停止的异常"""
try:
if self.translator:
self.translator.stop()
except Exception as e:
if "has stopped" not in str(e):
session_logger.warning(
self.device_id,
self.session_id,
f"停止translator时出现非标准错误: {str(e)}",
)
return None
async def transcribe_stream(asr: AliyunASR, stream_generator):
"""
流式转录音频数据
stream_generator: 生成器函数,每次产生音频块(bytes)
返回: 完整转录文本
"""
try:
# 连接ASR服务
await asr.connect()
if not asr.connected:
raise ConnectionError("无法连接到ASR服务")
# 处理音频流
for audio_chunk in stream_generator:
# 检查中断
# if interrupt_handler.is_interrupted((asr.device_id, asr.session_id)):
# session_logger.info(
# asr.device_id,
# asr.session_id,
# "检测到中断,停止音频发送"
# )
# return ""
await asr.send_audio(audio_chunk)
# 可选:获取中间结果
intermediate = await asr.receive_results()
if intermediate:
yield intermediate
# 发送结束信号并获取最终结果
await asr.send_end()
# return await asr.receive_results()
except Exception as e:
# session_logger.error(
# asr.device_id,
# asr.session_id,
# f"流式转录出错: {str(e)}"
# )
print(e)
yield ""
finally:
await asr.close()
# def audio_chunk_generator1(file_path, chunk_duration=100):
# """
# 生成器函数:从音频文件产生指定时长的块
# chunk_duration - 毫秒
# """
# audio = AudioProcessor.load_audio(file_path)
# chunk_size = int(16000 * 2 * (chunk_duration / 1000)) # 16kHz * 2字节 * 时长
# # 获取PCM数据
# pcm_data = AudioProcessor.audio_to_pcm(audio)
# # 分块产生
# for i in range(0, len(pcm_data), chunk_size):
# yield pcm_data[i:i+chunk_size]
# # 模拟实时延迟
# # await asyncio.sleep(chunk_duration / 1000)
def audio_chunk_generator(file_path, chunk_duration=100):
"""
生成器函数:从音频文件产生指定时长的块
chunk_duration - 毫秒
特点直接从AudioSegment对象生成音频块无需额外PCM转换
"""
# 加载音频文件
# audio = AudioProcessor.load_audio(file_path)
with open(file_path, "rb") as f:
audio = f.read()
# 计算每秒的字节数 (sample_rate * sample_width * channels)
# bytes_per_second = audio.frame_rate * (audio.sample_width) * audio.channels
# 计算每个块的字节数 (按时间计算而非固定值)
chunk_size_bytes = 320 #int(bytes_per_second * (chunk_duration / 1000.0))
# 计算块的数量
total_bytes = len(audio)
num_chunks = math.ceil(total_bytes / chunk_size_bytes)
# 分块产生音频数据
for i in range(num_chunks):
# 计算当前块的起始和结束位置
start_byte = i * chunk_size_bytes
end_byte = min((i + 1) * chunk_size_bytes, total_bytes)
# 提取原始字节数据
chunk_data = audio[start_byte:end_byte]
yield chunk_data
def create_websocket_packet(device_id: str, session_id: str, sequence_number: int, packet_type: int, audio_data: bytes) -> bytes:
"""
创建符合WebSocket接口格式的数据包
"""
# 设备ID以null结尾
device_id_bytes = device_id.encode('ascii') + b'\x00'
# 会话ID固定33字节不足补null
session_id_bytes = session_id.encode('ascii')
session_id_bytes = session_id_bytes.ljust(33, b'\x00')[:33]
# 序列号4字节小端
seq_bytes = struct.pack('<I', sequence_number)
# 包类型1字节
type_bytes = struct.pack('<B', packet_type)
# 音频数据大小4字节小端
size_bytes = struct.pack('<I', len(audio_data))
# 组合数据包
packet = device_id_bytes + session_id_bytes + seq_bytes + type_bytes + size_bytes + audio_data
return packet
async def send_audio_to_websocket(file_path: str, websocket_url: str = "ws://localhost:8000/ws"):
"""
将音频文件通过WebSocket发送到指定接口并接收打印返回结果
"""
device_id = "TalkingQ_10003B58862C"
session_id = "test-session-001"
# 加载音频文件
audio = AudioProcessor.load_audio(file_path)
# 转换为16kHz单声道PCM格式
pcm_audio = audio.set_frame_rate(16000).set_channels(1)
pcm_data = pcm_audio.raw_data
# 分块发送音频
chunk_size = 320 # 约100ms的音频数据
sequence_number = 0
async with websockets.connect(websocket_url) as websocket:
print(f"连接到WebSocket: {websocket_url}")
# 发送开始包
start_packet = create_websocket_packet(device_id, session_id, sequence_number, 1, b"")
await websocket.send(start_packet)
sequence_number += 1
print(f"发送开始包: 设备ID={device_id}, 会话ID={session_id}")
# 接收开始响应
try:
response = await asyncio.wait_for(websocket.recv(), timeout=5.0)
print(f"开始响应: {response}")
except asyncio.TimeoutError:
print("等待开始响应超时")
# 发送音频数据
for i in range(0, len(pcm_data), chunk_size):
if i>=3:
break
chunk = pcm_data[i:i+chunk_size]
audio_packet = create_websocket_packet(device_id, session_id, sequence_number, 0, chunk)
await websocket.send(audio_packet)
# 接收中间结果
try:
response = await asyncio.wait_for(websocket.recv(), timeout=1.0)
print(f"中间结果: {response}")
except asyncio.TimeoutError:
pass # 没有中间结果时继续
sequence_number += 1
# 添加延迟模拟实时流
await asyncio.sleep(0.1)
# 发送结束包
end_packet = create_websocket_packet(device_id, session_id, sequence_number, 2, b"")
await websocket.send(end_packet)
print("发送结束包")
# 接收最终结果
try:
response = await asyncio.wait_for(websocket.recv(), timeout=10.0)
print(f"最终结果: {response}")
except asyncio.TimeoutError:
print("等待最终结果超时")
print(f"音频文件 {file_path} 发送完成")
async def stream_large_audio():
asr = AliyunASR()
asr.device_id = "stream-device"
asr.session_id = "large-file-session"
# 创建音频流生成器
generator = audio_chunk_generator("./test/output(1).mp3", chunk_duration=500)
# generator1 = audio_chunk_generator("./test/16k.amr", chunk_duration=500)
# 转录MP3文件
print("开始转录MP3文件...")
async for result in transcribe_stream(asr, generator):
print(f"MP3中间结果: {result}")
async def send_audio_with_continuous_receive(file_path: str, websocket_url: str = "ws://118.25.13.180:8080/ws"):
"""
发送音频并持续接收所有返回结果
"""
device_id = "TalkingQ_10003B58862C"
session_id = "1000000189ABCDEFGHIJKLMNOPQRSTUV"
# 加载音频文件
# audio = AudioProcessor.load_audio(file_path)
with open(file_path, "rb") as f:
pcm_data = f.read()
# pcm_audio = audio.set_frame_rate(16000).set_channels(1)
# pcm_data = pcm_audio.raw_data
# pcm_data = AudioProcessor.audio_to_pcm(pcm_audio)
chunk_size = 320
sequence_number = 0
print(len(pcm_data))
async with websockets.connect(websocket_url) as websocket:
print(f"连接到WebSocket: {websocket_url}")
# 启动接收任务
receive_task = asyncio.create_task(receive_messages(websocket))
try:
# 发送开始包
start_packet = create_websocket_packet(device_id, session_id, sequence_number, 1, b"")
await websocket.send(start_packet)
sequence_number += 1
print("发送开始包")
# 发送音频数据
for i in range(0, len(pcm_data), chunk_size):
# if i>=1:
# break
chunk = pcm_data[i:i+chunk_size]
audio_packet = create_websocket_packet(device_id, session_id, sequence_number, 0, chunk)
await websocket.send(audio_packet)
sequence_number += 1
await asyncio.sleep(1)
# 发送结束包
end_packet = create_websocket_packet(device_id, session_id, sequence_number, 2, b"")
await websocket.send(end_packet)
print("发送结束包")
# 等待接收任务完成
await receive_task
except Exception as e:
print(f"发送过程中出错: {e}")
receive_task.cancel()
async def receive_messages(websocket):
"""
持续接收WebSocket消息
"""
try:
while True:
try:
message = await asyncio.wait_for(websocket.recv(), timeout=60.0)
print(f"收到消息: {message}")
except asyncio.TimeoutError:
print("接收超时,停止监听")
break
except websockets.exceptions.ConnectionClosed:
print("WebSocket连接已关闭")
break
except Exception as e:
print(f"接收消息时出错: {e}")
if __name__ == "__main__":
# 运行WebSocket测试 - 简单版本
# asyncio.run(send_audio_to_websocket("./test/output(1).mp3"))
# 运行WebSocket测试 - 持续接收版本
asyncio.run(send_audio_with_continuous_receive("./test/output(1).mp3"))
# 或者运行原来的测试
# asyncio.run(stream_large_audio())

View File

@@ -0,0 +1,55 @@
import requests
url = "https://api.minimaxi.com/v1/t2a_v2"
offline="你的好朋友现在不在哦"
welcome="你的好朋友已经在线,现在开始对话吧"
payload = {
"model": "speech-02-hd",
"text": welcome,
"voice_setting": {
"voice_id": "Chinese (Mandarin)_Cute_Spirit",
"speed": float(1.0),
"vol": 1.0,
"pitch": 0,
"emotion": 'happy',
},
"audio_setting": {
"sample_rate": 16000,
"bitrate": 32000,
"format": "mp3",
"channel": 1
},
"language_boost": "zh"
}
headers = {
"Content-Type": "application/json",
"Authorization": "Bearer eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJHcm91cE5hbWUiOiLovbvoiJ_mmbrlkK_vvIjmna3lt57vvInnp5HmioDmnInpmZDlhazlj7giLCJVc2VyTmFtZSI6Im1veSIsIkFjY291bnQiOiJtb3lAMTkxNTI5MjQxMDAyNDMwMDYzMCIsIlN1YmplY3RJRCI6IjE5MTY3OTA4MTQyNTY2NjQ2MDAiLCJQaG9uZSI6IiIsIkdyb3VwSUQiOiIxOTE1MjkyNDEwMDI0MzAwNjMwIiwiUGFnZU5hbWUiOiIiLCJNYWlsIjoiIiwiQ3JlYXRlVGltZSI6IjIwMjUtMDUtMDYgMTE6Mzg6MTMiLCJUb2tlblR5cGUiOjEsImlzcyI6Im1pbmltYXgifQ.Gw48hGemgBRA7YzjsWz5N2Vun7XRKyBXAKQLAnRZY6FdQQfDn__ZEUCMNVxMKcRns60-FrH-3Xp-nH-8nmQ-V67XtJ_JnBS4TH0NKDKt-vzj0xHmEWsMOEcwE24fOh1HB2U2o6teeLWJT_0fc2wRsdMD84NcRe0DSZ-Yi-et3_fbO8fmc6MTXGPkGQvYzS9k21Lm7A6rRUonL6IbTpqHSYuvA7bEV-6925gLxhwjSJ7q3r8KltKP4daGm2sIXhjr0mftR5t-NJs7pz-IWqoW5Lsaf3KgsIXTJn1wKwQtfo5L9oQxR7v8JifAOqJ_XPei5fXBYqIe5AoNATz5oKEo_A"
}
response = requests.post(url, json=payload, headers=headers)
print(response.json())
audio_data = response.json().get('data', {}).get('audio', '')
# print(audio_data)
if not audio_data:
raise Exception(f"Failed to get audio data from response")
# hex->bytes
output_file = "welcome.mp3"
audio_bytes = bytes.fromhex(audio_data)
with open(output_file, "wb") as f:
f.write(audio_bytes)
# import requests
# url = "https://api.minimaxi.com/v1/get_voice"
# payload = { "voice_type": "all" }
# headers = {
# "Content-Type": "application/json",
# "Authorization": "Bearer eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJHcm91cE5hbWUiOiLovbvoiJ_mmbrlkK_vvIjmna3lt57vvInnp5HmioDmnInpmZDlhazlj7giLCJVc2VyTmFtZSI6Im1veSIsIkFjY291bnQiOiJtb3lAMTkxNTI5MjQxMDAyNDMwMDYzMCIsIlN1YmplY3RJRCI6IjE5MTY3OTA4MTQyNTY2NjQ2MDAiLCJQaG9uZSI6IiIsIkdyb3VwSUQiOiIxOTE1MjkyNDEwMDI0MzAwNjMwIiwiUGFnZU5hbWUiOiIiLCJNYWlsIjoiIiwiQ3JlYXRlVGltZSI6IjIwMjUtMDUtMDYgMTE6Mzg6MTMiLCJUb2tlblR5cGUiOjEsImlzcyI6Im1pbmltYXgifQ.Gw48hGemgBRA7YzjsWz5N2Vun7XRKyBXAKQLAnRZY6FdQQfDn__ZEUCMNVxMKcRns60-FrH-3Xp-nH-8nmQ-V67XtJ_JnBS4TH0NKDKt-vzj0xHmEWsMOEcwE24fOh1HB2U2o6teeLWJT_0fc2wRsdMD84NcRe0DSZ-Yi-et3_fbO8fmc6MTXGPkGQvYzS9k21Lm7A6rRUonL6IbTpqHSYuvA7bEV-6925gLxhwjSJ7q3r8KltKP4daGm2sIXhjr0mftR5t-NJs7pz-IWqoW5Lsaf3KgsIXTJn1wKwQtfo5L9oQxR7v8JifAOqJ_XPei5fXBYqIe5AoNATz5oKEo_A"
# }
# response = requests.post(url, json=payload, headers=headers)
# print(response.json())

View File

@@ -0,0 +1,118 @@
import noisereduce as nr
import soundfile as sf
import librosa
import numpy as np
import os
def normalize_audio_volume(audio, original_rms=None, target_dbfs=-20):
"""
对音频进行音量归一化
参数:
audio: 音频数据
original_rms: 原始音频的RMS值可选
target_dbfs: 目标音量级别dBFS
返回:
归一化后的音频
"""
# 计算当前音频的RMS
current_rms = np.sqrt(np.mean(audio**2))
# 方法1恢复到原始音量水平
if original_rms is not None and original_rms > 0:
scale_factor = original_rms / current_rms
normalized_audio = audio * scale_factor
print(f"音量恢复: 缩放因子 {scale_factor:.2f}")
# 方法2基于目标dBFS的音量归一化
else:
# 计算当前音频的dBFS
current_dbfs = 20 * np.log10(current_rms / (np.max(np.abs(audio)) + 1e-8))
# 计算需要的增益
gain_db = target_dbfs - current_dbfs
gain_linear = 10 ** (gain_db / 20)
normalized_audio = audio * gain_linear
print(f"音量归一化: 增益 {gain_db:.1f} dB, 缩放因子 {gain_linear:.2f}")
# 防止削波clipping
max_val = np.max(np.abs(normalized_audio))
if max_val > 0.95: # 如果接近最大值1.0
normalized_audio = normalized_audio * 0.95 / max_val
print("已应用防削波处理")
return normalized_audio
def reduce_background_noise(audio_path, output_path=None, use_default_noise=True,
noise_path='assets/audio/noise_sample.wav', noise_start=0.1, noise_end=0.5,
normalize_volume=True, target_dbfs=-20):
"""
使用noisereduce库去除背景噪声
参数:
audio_path: 音频文件路径
output_path: 输出文件路径(可选)
use_default_noise: 是否使用默认噪声样本
noise_path: 默认噪声样本路径
noise_start: 噪声样本开始时间仅当use_default_noise=False时使用
noise_end: 噪声样本结束时间仅当use_default_noise=False时使用
"""
# 加载音频
y, sr = librosa.load(audio_path, sr=None)
print(f"音频信息: 时长 {len(y)/sr:.2f}秒, 采样率 {sr}Hz")
# 提取或加载噪声样本
if use_default_noise:
try:
noise_path = os.path.join("assets", "audio", "noise_sample.wav")
noise_clip, noise_sr = librosa.load(noise_path, sr=None)
# 检查采样率是否匹配
if noise_sr != sr:
print(f"警告: 噪声样本采样率({noise_sr}Hz)与音频采样率({sr}Hz)不匹配")
# 重采样噪声样本以匹配音频采样率
noise_clip = librosa.resample(noise_clip, orig_sr=noise_sr, target_sr=sr)
print("已对噪声样本进行重采样以匹配音频采样率")
except FileNotFoundError as e:
print(f"错误: {e}")
print("将使用当前音频提取噪声样本")
use_default_noise = False
if not use_default_noise:
# 从当前音频提取噪声样本
noise_start_sample = int(noise_start * sr)
noise_end_sample = int(noise_end * sr)
noise_clip = y[noise_start_sample:noise_end_sample]
print(f"使用当前音频的 {noise_end-noise_start:.2f}秒 噪声样本进行降噪")
else:
print(f"使用默认噪声样本进行降噪,噪声长度: {len(noise_clip)/sr:.2f}")
# 应用降噪
reduced_noise = nr.reduce_noise(
y=y,
sr=sr,
y_noise=noise_clip,
prop_decrease=0.95, # 降噪比例
n_fft=1024,
win_length=1024,
hop_length=256,
n_std_thresh_stationary=1.5,
stationary=True
)
# 记录原始音频的音量RMS
original_rms = np.sqrt(np.mean(y**2))
print(f"原始音频RMS: {original_rms:.4f}")
# 音量归一化处理
if normalize_volume:
normalize_audio_volume(reduced_noise, original_rms, target_dbfs)
# 保存结果
if output_path:
sf.write(output_path, reduced_noise, sr)
print(f"降噪后的音频已保存至: {output_path}")
if __name__ == "__main__":
reduce_background_noise('test/TalkingQ_XQSN00001004_eba844b5.mp3', 'test/denoised2.mp3')