Merge pull request #1074 from xinnan-tech/test-pr

增加server通用secret过滤器
This commit is contained in:
欣南科技
2025-04-30 15:11:09 +08:00
committed by GitHub
24 changed files with 347 additions and 200 deletions
+1 -2
View File
@@ -1,7 +1,7 @@
import asyncio
import sys
import signal
from config.settings import load_config, check_config_file
from config.settings import load_config
from core.websocket_server import WebSocketServer
from core.ota_server import SimpleOtaServer
from core.utils.util import check_ffmpeg_installed
@@ -31,7 +31,6 @@ async def wait_for_exit():
async def main():
check_config_file()
check_ffmpeg_installed()
config = load_config()
+8 -5
View File
@@ -1,7 +1,7 @@
# 如果您是一名开发者,建议阅读以下内容。如果不是开发者,可以忽略这部分内容。
# 在开发中,在项目根目录创建data目录,将【config.yaml】复制一份,改成【.config.yaml】,放进data目录中
# 系统会优先读取【data/.config.yaml】文件的配置。
# 这样做,可以避免在提交代码的时候,错误地提交密钥信息,保护您的密钥安全。
# 在开发中,请在项目根目录创建data目录,然后在data目录创建名称为【.config.yaml】的空文件
# 然后你想修改覆盖修改什么配置,就修改【.config.yaml】文件,而不是修改【config.yaml】文件
# 系统会优先读取【data/.config.yaml】文件的配置,如果【.config.yaml】文件里的配置不存在,系统会自动去读取【config.yaml】文件的配置
# 这样做,可以最简化配置,保护您的密钥安全。
# #####################################################################################
# #############################以下是服务器基本运行配置####################################
@@ -158,7 +158,7 @@ selected_module:
# 不想开通意图识别,就设置成:nointent
# 意图识别可使用intent_llm。优点:通用性强,缺点:增加串行前置意图识别模块,会增加处理时间,这个意图识别暂时不支持控制音量大小等iot操作
# 意图识别可使用function_call,缺点:需要所选择的LLM支持function_call,优点:按需调用工具、速度快,理论上能全部操作所有iot指令
# 默认免费的ChatGLMLLM就已经支持function_call,但是如果像追求稳定建议把LLM设置成:DoubaoLLM,使用的具体model_name是:doubao-pro-32k-functioncall-241028
# 默认免费的ChatGLMLLM就已经支持function_call,但是如果像追求稳定建议把LLM设置成:DoubaoLLM,使用的具体model_name是:doubao-1-5-pro-32k-250115
Intent: function_call
# 意图识别,是用于理解用户意图的模块,例如:播放音乐
@@ -397,6 +397,9 @@ TTS:
appid: 你的火山引擎语音合成服务appid
access_token: 你的火山引擎语音合成服务access_token
cluster: volcano_tts
speed_ratio: 1.0
volume_ratio: 1.0
pitch_ratio: 1.0
CosyVoiceSiliconflow:
type: siliconflow
# 硅基流动TTS
+42 -21
View File
@@ -1,6 +1,7 @@
import os
import argparse
import yaml
from collections.abc import Mapping
from config.manage_api_client import init_service, get_server_config, get_agent_models
@@ -25,35 +26,24 @@ def load_config():
if _config_cache is not None:
return _config_cache
parser = argparse.ArgumentParser(description="Server configuration")
config_file = get_config_file()
default_config_path = get_project_dir() + "config.yaml"
custom_config_path = get_project_dir() + "data/.config.yaml"
parser.add_argument("--config_path", type=str, default=config_file)
args = parser.parse_args()
config = read_config(args.config_path)
if config.get("manager-api", {}).get("url"):
config = get_config_from_api(config)
# 加载默认配置
default_config = read_config(default_config_path)
custom_config = read_config(custom_config_path)
if custom_config.get("manager-api", {}).get("url"):
config = get_config_from_api(custom_config)
else:
# 合并配置
config = merge_configs(default_config, custom_config)
# 初始化目录
ensure_directories(config)
_config_cache = config
return config
def get_config_file():
"""获取配置文件路径,优先使用私有配置文件(若存在)。
Returns:
str: 配置文件路径(相对路径或默认路径)
"""
default_config_file = "config.yaml"
config_file = default_config_file
if os.path.exists(get_project_dir() + "data/." + default_config_file):
config_file = "data/." + default_config_file
return config_file
def get_config_from_api(config):
"""从Java API获取配置"""
# 初始化API客户端
@@ -115,3 +105,34 @@ def ensure_directories(config):
os.makedirs(dir_path, exist_ok=True)
except PermissionError:
print(f"警告:无法创建目录 {dir_path},请检查写入权限")
def merge_configs(default_config, custom_config):
"""
递归合并配置,custom_config优先级更高
Args:
default_config: 默认配置
custom_config: 用户自定义配置
Returns:
合并后的配置
"""
if not isinstance(default_config, Mapping) or not isinstance(
custom_config, Mapping
):
return custom_config
merged = dict(default_config)
for key, value in custom_config.items():
if (
key in merged
and isinstance(merged[key], Mapping)
and isinstance(value, Mapping)
):
merged[key] = merge_configs(merged[key], value)
else:
merged[key] = value
return merged
+2
View File
@@ -2,6 +2,7 @@ import os
import sys
from loguru import logger
from config.config_loader import load_config
from config.settings import check_config_file
SERVER_VERSION = "0.3.13"
@@ -32,6 +33,7 @@ def formatter(record):
def setup_logging():
check_config_file()
"""从配置文件中读取日志配置,并设置日志输出格式和级别"""
config = load_config()
log_config = config["log"]
@@ -53,6 +53,7 @@ class ManageApiClient:
headers={
"User-Agent": f"PythonClient/2.0 (PID:{os.getpid()})",
"Accept": "application/json",
"Authorization": "Bearer " + cls._secret
},
timeout=cls.config.get("timeout", 30), # 默认超时时间30秒
)
@@ -126,7 +127,7 @@ class ManageApiClient:
def get_server_config() -> Optional[Dict]:
"""获取服务器基础配置"""
return ManageApiClient._instance._execute_request(
"POST", "/config/server-base", json={"secret": ManageApiClient._secret}
"POST", "/config/server-base"
)
@@ -138,7 +139,6 @@ def get_agent_models(
"POST",
"/config/agent-models",
json={
"secret": ManageApiClient._secret,
"macAddress": mac_address,
"clientId": client_id,
"selectedModule": selected_module,
+19 -50
View File
@@ -1,64 +1,33 @@
import os
from collections.abc import Mapping
from config.config_loader import read_config, get_project_dir, load_config
default_config_file = "config.yaml"
def find_missing_keys(new_config, old_config, parent_key=""):
"""
递归查找缺失的配置项
返回格式:[缺失配置路径]
"""
missing_keys = []
if not isinstance(new_config, Mapping):
return missing_keys
for key, value in new_config.items():
# 构建当前配置路径
full_path = f"{parent_key}.{key}" if parent_key else key
# 检查键是否存在
if key not in old_config:
missing_keys.append(full_path)
continue
# 递归检查嵌套字典
if isinstance(value, Mapping):
sub_missing = find_missing_keys(
value, old_config[key], parent_key=full_path
)
missing_keys.extend(sub_missing)
return missing_keys
config_file_valid = False
def check_config_file():
old_config_file = get_project_dir() + "data/." + default_config_file
if not os.path.exists(old_config_file):
global config_file_valid
if config_file_valid:
return
old_config = load_config()
new_config = read_config(get_project_dir() + default_config_file)
# 查找缺失的配置项
missing_keys = find_missing_keys(new_config, old_config)
read_config_from_api = old_config.get("read_config_from_api", False)
if read_config_from_api:
old_config_origin = read_config(old_config_file)
"""
简化的配置检查,仅提示用户配置文件的使用情况
"""
custom_config_file = get_project_dir() + "data/." + default_config_file
if not os.path.exists(custom_config_file):
raise FileNotFoundError(
"找不到data/.config.yaml文件,请按教程确认该配置文件是否存在"
)
# 检查是否从API读取配置
config = load_config()
if config.get("read_config_from_api", False):
print("从API读取配置")
old_config_origin = read_config(custom_config_file)
if old_config_origin.get("selected_module") is not None:
missing_keys_str = "\n".join(f"- {key}" for key in missing_keys)
error_msg = "您的配置文件好像既包含智控台的配置又包含本地配置:\n"
error_msg += "\n建议您:\n"
error_msg += "1、将根目录的config_from_api.yaml文件复制到data下,重命名为.config.yaml\n"
error_msg += "2、按教程配置好接口地址和密钥\n"
raise ValueError(error_msg)
return
if missing_keys:
missing_keys_str = "\n".join(f"- {key}" for key in missing_keys)
error_msg = "您的配置文件太旧了,缺少了:\n"
error_msg += missing_keys_str
error_msg += "\n建议您:\n"
error_msg += "1、备份data/.config.yaml文件\n"
error_msg += "2、将根目录的config.yaml文件复制到data下,重命名为.config.yaml\n"
error_msg += "3、将密钥逐个复制到新的配置文件中\n"
raise ValueError(error_msg)
config_file_valid = True
@@ -53,6 +53,8 @@ async def handleTextMessage(conn, message):
# 如果是唤醒词,且关闭了唤醒词回复,就不用回答
await send_stt_message(conn, text)
await send_tts_message(conn, "stop", None)
elif is_wakeup_words:
await startToChat(conn, "嘿,你好呀")
else:
# 否则需要LLM对文字内容进行答复
await startToChat(conn, text)
@@ -27,6 +27,15 @@ class TTSProvider(TTSProviderBase):
else:
self.voice = config.get("voice")
# 处理空字符串的情况
speed_ratio = config.get("speed_ratio", "1.0")
volume_ratio = config.get("volume_ratio", "1.0")
pitch_ratio = config.get("pitch_ratio", "1.0")
self.speed_ratio = float(speed_ratio) if speed_ratio else 1.0
self.volume_ratio = float(volume_ratio) if volume_ratio else 1.0
self.pitch_ratio = float(pitch_ratio) if pitch_ratio else 1.0
self.api_url = config.get("api_url")
self.authorization = config.get("authorization")
self.header = {"Authorization": f"{self.authorization}{self.access_token}"}
@@ -49,9 +58,9 @@ class TTSProvider(TTSProviderBase):
"audio": {
"voice_type": self.voice,
"encoding": "wav",
"speed_ratio": 1.0,
"volume_ratio": 1.0,
"pitch_ratio": 1.0,
"speed_ratio": self.speed_ratio,
"volume_ratio": self.volume_ratio,
"pitch_ratio": self.pitch_ratio,
},
"request": {
"reqid": str(uuid.uuid4()),
@@ -89,18 +89,35 @@ class TTSProvider(TTSProviderBase):
self.reference_audio = parse_string_to_list(config.get("reference_audio"))
self.reference_text = parse_string_to_list(config.get("reference_text"))
self.format = config.get("format", "wav")
self.channels = int(config.get("channels", 1))
self.rate = int(config.get("rate", 44100))
self.api_key = config.get("api_key", "YOUR_API_KEY")
have_key = check_model_key("FishSpeech TTS", self.api_key)
if not have_key:
return
self.normalize = config.get("normalize", True)
self.max_new_tokens = int(config.get("max_new_tokens", 1024))
self.chunk_length = int(config.get("chunk_length", 200))
self.top_p = float(config.get("top_p", 0.7))
self.repetition_penalty = float(config.get("repetition_penalty", 1.2))
self.temperature = float(config.get("temperature", 0.7))
# 处理空字符串的情况
channels = config.get("channels", "1")
rate = config.get("rate", "44100")
max_new_tokens = config.get("max_new_tokens", "1024")
chunk_length = config.get("chunk_length", "200")
self.channels = int(channels) if channels else 1
self.rate = int(rate) if rate else 44100
self.max_new_tokens = int(max_new_tokens) if max_new_tokens else 1024
self.chunk_length = int(chunk_length) if chunk_length else 200
# 处理空字符串的情况
top_p = config.get("top_p", "0.7")
temperature = config.get("temperature", "0.7")
repetition_penalty = config.get("repetition_penalty", "1.2")
self.top_p = float(top_p) if top_p else 0.7
self.temperature = float(temperature) if temperature else 0.7
self.repetition_penalty = (
float(repetition_penalty) if repetition_penalty else 1.2
)
self.streaming = str(config.get("streaming", False)).lower() in (
"true",
"1",
@@ -20,12 +20,29 @@ class TTSProvider(TTSProviderBase):
self.ref_audio_path = config.get("ref_audio_path")
self.prompt_text = config.get("prompt_text")
self.prompt_lang = config.get("prompt_lang", "zh")
self.top_k = int(config.get("top_k", 5))
self.top_p = float(config.get("top_p", 1))
self.temperature = float(config.get("temperature", 1))
# 处理空字符串的情况
top_k = config.get("top_k", "5")
top_p = config.get("top_p", "1")
temperature = config.get("temperature", "1")
batch_threshold = config.get("batch_threshold", "0.75")
batch_size = config.get("batch_size", "1")
speed_factor = config.get("speed_factor", "1.0")
seed = config.get("seed", "-1")
repetition_penalty = config.get("repetition_penalty", "1.35")
self.top_k = int(top_k) if top_k else 5
self.top_p = float(top_p) if top_p else 1
self.temperature = float(temperature) if temperature else 1
self.batch_threshold = float(batch_threshold) if batch_threshold else 0.75
self.batch_size = int(batch_size) if batch_size else 1
self.speed_factor = float(speed_factor) if speed_factor else 1.0
self.seed = int(seed) if seed else -1
self.repetition_penalty = (
float(repetition_penalty) if repetition_penalty else 1.35
)
self.text_split_method = config.get("text_split_method", "cut0")
self.batch_size = int(config.get("batch_size", 1))
self.batch_threshold = float(config.get("batch_threshold", 0.75))
self.split_bucket = str(config.get("split_bucket", True)).lower() in (
"true",
@@ -37,19 +54,19 @@ class TTSProvider(TTSProviderBase):
"1",
"yes",
)
self.speed_factor = float(config.get("speed_factor", 1.0))
self.streaming_mode = str(config.get("streaming_mode", False)).lower() in (
"true",
"1",
"yes",
)
self.seed = int(config.get("seed", -1))
self.parallel_infer = str(config.get("parallel_infer", True)).lower() in (
"true",
"1",
"yes",
)
self.repetition_penalty = float(config.get("repetition_penalty", 1.35))
self.aux_ref_audio_paths = parse_string_to_list(
config.get("aux_ref_audio_paths")
)
@@ -18,13 +18,22 @@ class TTSProvider(TTSProviderBase):
self.prompt_text = config.get("prompt_text")
self.prompt_language = config.get("prompt_language")
self.text_language = config.get("text_language", "audo")
self.top_k = int(config.get("top_k", 15))
self.top_p = float(config.get("top_p", 1.0))
self.temperature = float(config.get("temperature", 1.0))
# 处理空字符串的情况
top_k = config.get("top_k", "15")
top_p = config.get("top_p", "1.0")
temperature = config.get("temperature", "1.0")
sample_steps = config.get("sample_steps", "32")
speed = config.get("speed", "1.0")
self.top_k = int(top_k) if top_k else 15
self.top_p = float(top_p) if top_p else 1.0
self.temperature = float(temperature) if temperature else 1.0
self.sample_steps = int(sample_steps) if sample_steps else 32
self.speed = float(speed) if speed else 1.0
self.cut_punc = config.get("cut_punc", "")
self.speed = float(config.get("speed", 1.0))
self.inp_refs = parse_string_to_list(config.get("inp_refs"))
self.sample_steps = int(config.get("sample_steps", 32))
self.if_sr = str(config.get("if_sr", False)).lower() in ("true", "1", "yes")
def generate_filename(self, extension=".wav"):
@@ -21,7 +21,11 @@ class TTSProvider(TTSProviderBase):
else:
self.voice = config.get("voice", "alloy")
self.response_format = "wav"
self.speed = float(config.get("speed", 1.0))
# 处理空字符串的情况
speed = config.get("speed", "1.0")
self.speed = float(speed) if speed else 1.0
self.output_file = config.get("output_dir", "tmp/")
check_model_key("TTS", self.api_key)
@@ -21,8 +21,15 @@ class VADProvider(VADProviderBase):
(get_speech_timestamps, _, _, _, _) = self.utils
self.decoder = opuslib_next.Decoder(16000, 1)
self.vad_threshold = float(config.get("threshold", 0.5))
self.silence_threshold_ms = int(config.get("min_silence_duration_ms", 1000))
# 处理空字符串的情况
threshold = config.get("threshold", "0.5")
min_silence_duration_ms = config.get("min_silence_duration_ms", "1000")
self.vad_threshold = float(threshold) if threshold else 0.5
self.silence_threshold_ms = (
int(min_silence_duration_ms) if min_silence_duration_ms else 1000
)
def is_vad(self, conn, opus_packet):
try:
+1 -1
View File
@@ -49,7 +49,7 @@ services:
- SPRING_DATA_REDIS_PORT=6379
volumes:
# 配置文件目录
- ./uploadfile:/app/uploadfile
- ./uploadfile:/uploadfile
xiaozhi-esp32-server-db:
image: mysql:latest