mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-27 01:23:55 +08:00
update:修复智控台下发配置布尔类型转换出错bug (#850)
* update:测试页面增加OTA地址 * update:兼容旧设备,无Client-Id的情况 * update:修复智控台下发配置布尔类型转换出错bug * update:修复智控台下发配置字符类型转换出错bug
This commit is contained in:
@@ -4,9 +4,17 @@ import json
|
|||||||
import re
|
import re
|
||||||
from core.providers.llm.base import LLMProviderBase
|
from core.providers.llm.base import LLMProviderBase
|
||||||
import os
|
import os
|
||||||
|
|
||||||
# official coze sdk for Python [cozepy](https://github.com/coze-dev/coze-py)
|
# official coze sdk for Python [cozepy](https://github.com/coze-dev/coze-py)
|
||||||
from cozepy import COZE_CN_BASE_URL
|
from cozepy import COZE_CN_BASE_URL
|
||||||
from cozepy import Coze, TokenAuth, Message, ChatStatus, MessageContentType, ChatEventType # noqa
|
from cozepy import (
|
||||||
|
Coze,
|
||||||
|
TokenAuth,
|
||||||
|
Message,
|
||||||
|
ChatStatus,
|
||||||
|
MessageContentType,
|
||||||
|
ChatEventType,
|
||||||
|
) # noqa
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -15,8 +23,8 @@ logger = setup_logging()
|
|||||||
class LLMProvider(LLMProviderBase):
|
class LLMProvider(LLMProviderBase):
|
||||||
def __init__(self, config):
|
def __init__(self, config):
|
||||||
self.personal_access_token = config.get("personal_access_token")
|
self.personal_access_token = config.get("personal_access_token")
|
||||||
self.bot_id = config.get("bot_id")
|
self.bot_id = str(config.get("bot_id"))
|
||||||
self.user_id = config.get("user_id")
|
self.user_id = str(config.get("user_id"))
|
||||||
self.session_conversation_map = {} # 存储session_id和conversation_id的映射
|
self.session_conversation_map = {} # 存储session_id和conversation_id的映射
|
||||||
|
|
||||||
def response(self, session_id, dialogue):
|
def response(self, session_id, dialogue):
|
||||||
@@ -30,10 +38,7 @@ class LLMProvider(LLMProviderBase):
|
|||||||
|
|
||||||
# 如果没有找到conversation_id,则创建新的对话
|
# 如果没有找到conversation_id,则创建新的对话
|
||||||
if not conversation_id:
|
if not conversation_id:
|
||||||
conversation = coze.conversations.create(
|
conversation = coze.conversations.create(messages=[])
|
||||||
messages=[
|
|
||||||
]
|
|
||||||
)
|
|
||||||
conversation_id = conversation.id
|
conversation_id = conversation.id
|
||||||
self.session_conversation_map[session_id] = conversation_id # 更新映射
|
self.session_conversation_map[session_id] = conversation_id # 更新映射
|
||||||
|
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ from pydantic import BaseModel, Field, conint, model_validator
|
|||||||
from typing_extensions import Annotated
|
from typing_extensions import Annotated
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
from core.utils.util import check_model_key
|
from core.utils.util import check_model_key, parse_string_to_list
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
|
|
||||||
@@ -86,8 +86,8 @@ class TTSProvider(TTSProviderBase):
|
|||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
|
|
||||||
self.reference_id = config.get("reference_id")
|
self.reference_id = config.get("reference_id")
|
||||||
self.reference_audio = config.get("reference_audio", [])
|
self.reference_audio = parse_string_to_list(config.get("reference_audio"))
|
||||||
self.reference_text = config.get("reference_text", [])
|
self.reference_text = parse_string_to_list(config.get("reference_text"))
|
||||||
self.format = config.get("format", "wav")
|
self.format = config.get("format", "wav")
|
||||||
self.channels = int(config.get("channels", 1))
|
self.channels = int(config.get("channels", 1))
|
||||||
self.rate = int(config.get("rate", 44100))
|
self.rate = int(config.get("rate", 44100))
|
||||||
@@ -101,9 +101,13 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.top_p = float(config.get("top_p", 0.7))
|
self.top_p = float(config.get("top_p", 0.7))
|
||||||
self.repetition_penalty = float(config.get("repetition_penalty", 1.2))
|
self.repetition_penalty = float(config.get("repetition_penalty", 1.2))
|
||||||
self.temperature = float(config.get("temperature", 0.7))
|
self.temperature = float(config.get("temperature", 0.7))
|
||||||
self.streaming = bool(config.get("streaming", False))
|
self.streaming = str(config.get("streaming", False)).lower() in (
|
||||||
|
"true",
|
||||||
|
"1",
|
||||||
|
"yes",
|
||||||
|
)
|
||||||
self.use_memory_cache = config.get("use_memory_cache", "on")
|
self.use_memory_cache = config.get("use_memory_cache", "on")
|
||||||
self.seed = config.get("seed")
|
self.seed = config.get("seed") or None
|
||||||
self.api_url = config.get("api_url", "http://127.0.0.1:8080/v1/tts")
|
self.api_url = config.get("api_url", "http://127.0.0.1:8080/v1/tts")
|
||||||
|
|
||||||
def generate_filename(self, extension=".wav"):
|
def generate_filename(self, extension=".wav"):
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import requests
|
|||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
|
from core.utils.util import parse_string_to_list
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -25,14 +26,33 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.text_split_method = config.get("text_split_method", "cut0")
|
self.text_split_method = config.get("text_split_method", "cut0")
|
||||||
self.batch_size = int(config.get("batch_size", 1))
|
self.batch_size = int(config.get("batch_size", 1))
|
||||||
self.batch_threshold = float(config.get("batch_threshold", 0.75))
|
self.batch_threshold = float(config.get("batch_threshold", 0.75))
|
||||||
self.split_bucket = bool(config.get("split_bucket", True))
|
|
||||||
self.return_fragment = bool(config.get("return_fragment", False))
|
self.split_bucket = str(config.get("split_bucket", True)).lower() in (
|
||||||
|
"true",
|
||||||
|
"1",
|
||||||
|
"yes",
|
||||||
|
)
|
||||||
|
self.return_fragment = str(config.get("return_fragment", False)).lower() in (
|
||||||
|
"true",
|
||||||
|
"1",
|
||||||
|
"yes",
|
||||||
|
)
|
||||||
self.speed_factor = float(config.get("speed_factor", 1.0))
|
self.speed_factor = float(config.get("speed_factor", 1.0))
|
||||||
self.streaming_mode = bool(config.get("streaming_mode", False))
|
self.streaming_mode = str(config.get("streaming_mode", False)).lower() in (
|
||||||
|
"true",
|
||||||
|
"1",
|
||||||
|
"yes",
|
||||||
|
)
|
||||||
self.seed = int(config.get("seed", -1))
|
self.seed = int(config.get("seed", -1))
|
||||||
self.parallel_infer = bool(config.get("parallel_infer", True))
|
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.repetition_penalty = float(config.get("repetition_penalty", 1.35))
|
||||||
self.aux_ref_audio_paths = config.get("aux_ref_audio_paths", [])
|
self.aux_ref_audio_paths = parse_string_to_list(
|
||||||
|
config.get("aux_ref_audio_paths")
|
||||||
|
)
|
||||||
|
|
||||||
def generate_filename(self, extension=".wav"):
|
def generate_filename(self, extension=".wav"):
|
||||||
return os.path.join(
|
return os.path.join(
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import requests
|
|||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
|
from core.utils.util import parse_string_to_list
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -22,9 +23,9 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.temperature = float(config.get("temperature", 1.0))
|
self.temperature = float(config.get("temperature", 1.0))
|
||||||
self.cut_punc = config.get("cut_punc", "")
|
self.cut_punc = config.get("cut_punc", "")
|
||||||
self.speed = float(config.get("speed", 1.0))
|
self.speed = float(config.get("speed", 1.0))
|
||||||
self.inp_refs = config.get("inp_refs", [])
|
self.inp_refs = parse_string_to_list(config.get("inp_refs"))
|
||||||
self.sample_steps = int(config.get("sample_steps", 32))
|
self.sample_steps = int(config.get("sample_steps", 32))
|
||||||
self.if_sr = bool(config.get("if_sr", False))
|
self.if_sr = str(config.get("if_sr", False)).lower() in ("true", "1", "yes")
|
||||||
|
|
||||||
def generate_filename(self, extension=".wav"):
|
def generate_filename(self, extension=".wav"):
|
||||||
return os.path.join(
|
return os.path.join(
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import json
|
|||||||
import requests
|
import requests
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
|
from core.utils.util import parse_string_to_list
|
||||||
|
|
||||||
|
|
||||||
class TTSProvider(TTSProviderBase):
|
class TTSProvider(TTSProviderBase):
|
||||||
@@ -40,7 +41,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
**config.get("pronunciation_dict", {}),
|
**config.get("pronunciation_dict", {}),
|
||||||
}
|
}
|
||||||
self.audio_setting = {**defult_audio_setting, **config.get("audio_setting", {})}
|
self.audio_setting = {**defult_audio_setting, **config.get("audio_setting", {})}
|
||||||
self.timber_weights = config.get("timber_weights", [])
|
self.timber_weights = parse_string_to_list(config.get("timber_weights"))
|
||||||
|
|
||||||
if self.voice_id:
|
if self.voice_id:
|
||||||
self.voice_setting["voice_id"] = self.voice_id
|
self.voice_setting["voice_id"] = self.voice_id
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.to_lang = config.get("to_lang")
|
self.to_lang = config.get("to_lang")
|
||||||
self.volume_change_dB = int(config.get("volume_change_dB", 0))
|
self.volume_change_dB = int(config.get("volume_change_dB", 0))
|
||||||
self.speed_factor = int(config.get("speed_factor", 1))
|
self.speed_factor = int(config.get("speed_factor", 1))
|
||||||
self.stream = bool(config.get("stream", False))
|
self.stream = str(config.get("stream", False)).lower() in ("true", "1", "yes")
|
||||||
self.output_file = config.get("output_dir")
|
self.output_file = config.get("output_dir")
|
||||||
self.pitch_factor = int(config.get("pitch_factor", 0))
|
self.pitch_factor = int(config.get("pitch_factor", 0))
|
||||||
self.format = config.get("format", "mp3")
|
self.format = config.get("format", "mp3")
|
||||||
|
|||||||
@@ -163,6 +163,24 @@ def check_model_key(modelType, modelKey):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def parse_string_to_list(value, separator=";"):
|
||||||
|
"""
|
||||||
|
将输入值转换为列表
|
||||||
|
Args:
|
||||||
|
value: 输入值,可以是 None、字符串或列表
|
||||||
|
separator: 分隔符,默认为分号
|
||||||
|
Returns:
|
||||||
|
list: 处理后的列表
|
||||||
|
"""
|
||||||
|
if value is None or value == "":
|
||||||
|
return []
|
||||||
|
elif isinstance(value, str):
|
||||||
|
return [item.strip() for item in value.split(separator) if item.strip()]
|
||||||
|
elif isinstance(value, list):
|
||||||
|
return value
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
def check_ffmpeg_installed():
|
def check_ffmpeg_installed():
|
||||||
ffmpeg_installed = False
|
ffmpeg_installed = False
|
||||||
try:
|
try:
|
||||||
@@ -231,7 +249,7 @@ def initialize_modules(
|
|||||||
modules["tts"] = tts.create_instance(
|
modules["tts"] = tts.create_instance(
|
||||||
tts_type,
|
tts_type,
|
||||||
config["TTS"][select_tts_module],
|
config["TTS"][select_tts_module],
|
||||||
bool(config.get("delete_audio", True)),
|
str(config.get("delete_audio", True)).lower() in ("true", "1", "yes"),
|
||||||
)
|
)
|
||||||
logger.bind(tag=TAG).info(f"初始化组件: tts成功 {select_tts_module}")
|
logger.bind(tag=TAG).info(f"初始化组件: tts成功 {select_tts_module}")
|
||||||
|
|
||||||
@@ -302,7 +320,7 @@ def initialize_modules(
|
|||||||
modules["asr"] = asr.create_instance(
|
modules["asr"] = asr.create_instance(
|
||||||
asr_type,
|
asr_type,
|
||||||
config["ASR"][select_asr_module],
|
config["ASR"][select_asr_module],
|
||||||
bool(config.get("delete_audio", True)),
|
str(config.get("delete_audio", True)).lower() in ("true", "1", "yes"),
|
||||||
)
|
)
|
||||||
logger.bind(tag=TAG).info(f"初始化组件: asr成功 {select_asr_module}")
|
logger.bind(tag=TAG).info(f"初始化组件: asr成功 {select_asr_module}")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user