mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-22 07:03:53 +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
|
||||
from core.providers.llm.base import LLMProviderBase
|
||||
import os
|
||||
|
||||
# official coze sdk for Python [cozepy](https://github.com/coze-dev/coze-py)
|
||||
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__
|
||||
logger = setup_logging()
|
||||
@@ -15,8 +23,8 @@ logger = setup_logging()
|
||||
class LLMProvider(LLMProviderBase):
|
||||
def __init__(self, config):
|
||||
self.personal_access_token = config.get("personal_access_token")
|
||||
self.bot_id = config.get("bot_id")
|
||||
self.user_id = config.get("user_id")
|
||||
self.bot_id = str(config.get("bot_id"))
|
||||
self.user_id = str(config.get("user_id"))
|
||||
self.session_conversation_map = {} # 存储session_id和conversation_id的映射
|
||||
|
||||
def response(self, session_id, dialogue):
|
||||
@@ -24,16 +32,13 @@ class LLMProvider(LLMProviderBase):
|
||||
coze_api_base = COZE_CN_BASE_URL
|
||||
|
||||
last_msg = next(m for m in reversed(dialogue) if m["role"] == "user")
|
||||
|
||||
|
||||
coze = Coze(auth=TokenAuth(token=coze_api_token), base_url=coze_api_base)
|
||||
conversation_id = self.session_conversation_map.get(session_id)
|
||||
|
||||
# 如果没有找到conversation_id,则创建新的对话
|
||||
if not conversation_id:
|
||||
conversation = coze.conversations.create(
|
||||
messages=[
|
||||
]
|
||||
)
|
||||
conversation = coze.conversations.create(messages=[])
|
||||
conversation_id = conversation.id
|
||||
self.session_conversation_map[session_id] = conversation_id # 更新映射
|
||||
|
||||
@@ -47,4 +52,4 @@ class LLMProvider(LLMProviderBase):
|
||||
):
|
||||
if event.event == ChatEventType.CONVERSATION_MESSAGE_DELTA:
|
||||
print(event.message.content, end="", flush=True)
|
||||
yield event.message.content
|
||||
yield event.message.content
|
||||
|
||||
@@ -8,7 +8,7 @@ from pydantic import BaseModel, Field, conint, model_validator
|
||||
from typing_extensions import Annotated
|
||||
from datetime import datetime
|
||||
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 config.logger import setup_logging
|
||||
|
||||
@@ -86,8 +86,8 @@ class TTSProvider(TTSProviderBase):
|
||||
super().__init__(config, delete_audio_file)
|
||||
|
||||
self.reference_id = config.get("reference_id")
|
||||
self.reference_audio = config.get("reference_audio", [])
|
||||
self.reference_text = config.get("reference_text", [])
|
||||
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))
|
||||
@@ -101,9 +101,13 @@ class TTSProvider(TTSProviderBase):
|
||||
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))
|
||||
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.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")
|
||||
|
||||
def generate_filename(self, extension=".wav"):
|
||||
|
||||
@@ -6,6 +6,7 @@ import requests
|
||||
from config.logger import setup_logging
|
||||
from datetime import datetime
|
||||
from core.providers.tts.base import TTSProviderBase
|
||||
from core.utils.util import parse_string_to_list
|
||||
|
||||
TAG = __name__
|
||||
logger = setup_logging()
|
||||
@@ -25,14 +26,33 @@ class TTSProvider(TTSProviderBase):
|
||||
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 = 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.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.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.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"):
|
||||
return os.path.join(
|
||||
|
||||
@@ -4,6 +4,7 @@ import requests
|
||||
from config.logger import setup_logging
|
||||
from datetime import datetime
|
||||
from core.providers.tts.base import TTSProviderBase
|
||||
from core.utils.util import parse_string_to_list
|
||||
|
||||
TAG = __name__
|
||||
logger = setup_logging()
|
||||
@@ -22,9 +23,9 @@ class TTSProvider(TTSProviderBase):
|
||||
self.temperature = float(config.get("temperature", 1.0))
|
||||
self.cut_punc = config.get("cut_punc", "")
|
||||
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.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"):
|
||||
return os.path.join(
|
||||
|
||||
@@ -4,6 +4,7 @@ import json
|
||||
import requests
|
||||
from datetime import datetime
|
||||
from core.providers.tts.base import TTSProviderBase
|
||||
from core.utils.util import parse_string_to_list
|
||||
|
||||
|
||||
class TTSProvider(TTSProviderBase):
|
||||
@@ -40,7 +41,7 @@ class TTSProvider(TTSProviderBase):
|
||||
**config.get("pronunciation_dict", {}),
|
||||
}
|
||||
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:
|
||||
self.voice_setting["voice_id"] = self.voice_id
|
||||
|
||||
@@ -22,7 +22,7 @@ class TTSProvider(TTSProviderBase):
|
||||
self.to_lang = config.get("to_lang")
|
||||
self.volume_change_dB = int(config.get("volume_change_dB", 0))
|
||||
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.pitch_factor = int(config.get("pitch_factor", 0))
|
||||
self.format = config.get("format", "mp3")
|
||||
|
||||
@@ -163,6 +163,24 @@ def check_model_key(modelType, modelKey):
|
||||
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():
|
||||
ffmpeg_installed = False
|
||||
try:
|
||||
@@ -231,7 +249,7 @@ def initialize_modules(
|
||||
modules["tts"] = tts.create_instance(
|
||||
tts_type,
|
||||
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}")
|
||||
|
||||
@@ -302,7 +320,7 @@ def initialize_modules(
|
||||
modules["asr"] = asr.create_instance(
|
||||
asr_type,
|
||||
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}")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user