diff --git a/main/xiaozhi-server/core/providers/llm/coze/coze.py b/main/xiaozhi-server/core/providers/llm/coze/coze.py index b1b14d91..41489e01 100644 --- a/main/xiaozhi-server/core/providers/llm/coze/coze.py +++ b/main/xiaozhi-server/core/providers/llm/coze/coze.py @@ -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 \ No newline at end of file + yield event.message.content diff --git a/main/xiaozhi-server/core/providers/tts/fishspeech.py b/main/xiaozhi-server/core/providers/tts/fishspeech.py index 6456961d..30fee3c9 100644 --- a/main/xiaozhi-server/core/providers/tts/fishspeech.py +++ b/main/xiaozhi-server/core/providers/tts/fishspeech.py @@ -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"): diff --git a/main/xiaozhi-server/core/providers/tts/gpt_sovits_v2.py b/main/xiaozhi-server/core/providers/tts/gpt_sovits_v2.py index 311ef5b0..ec6c2053 100644 --- a/main/xiaozhi-server/core/providers/tts/gpt_sovits_v2.py +++ b/main/xiaozhi-server/core/providers/tts/gpt_sovits_v2.py @@ -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( diff --git a/main/xiaozhi-server/core/providers/tts/gpt_sovits_v3.py b/main/xiaozhi-server/core/providers/tts/gpt_sovits_v3.py index 25c52b84..b2746acf 100644 --- a/main/xiaozhi-server/core/providers/tts/gpt_sovits_v3.py +++ b/main/xiaozhi-server/core/providers/tts/gpt_sovits_v3.py @@ -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( diff --git a/main/xiaozhi-server/core/providers/tts/minimax.py b/main/xiaozhi-server/core/providers/tts/minimax.py index 2b63ea7e..dd406b64 100644 --- a/main/xiaozhi-server/core/providers/tts/minimax.py +++ b/main/xiaozhi-server/core/providers/tts/minimax.py @@ -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 diff --git a/main/xiaozhi-server/core/providers/tts/ttson.py b/main/xiaozhi-server/core/providers/tts/ttson.py index a6401141..c2b78557 100644 --- a/main/xiaozhi-server/core/providers/tts/ttson.py +++ b/main/xiaozhi-server/core/providers/tts/ttson.py @@ -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") diff --git a/main/xiaozhi-server/core/utils/util.py b/main/xiaozhi-server/core/utils/util.py index 238bf5cb..f404a94c 100644 --- a/main/xiaozhi-server/core/utils/util.py +++ b/main/xiaozhi-server/core/utils/util.py @@ -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}")