update:修复智控台下发配置布尔类型转换出错bug (#850)

* update:测试页面增加OTA地址

* update:兼容旧设备,无Client-Id的情况

* update:修复智控台下发配置布尔类型转换出错bug

* update:修复智控台下发配置字符类型转换出错bug
This commit is contained in:
hrz
2025-04-16 22:55:13 +08:00
committed by GitHub
parent bfdfa44edd
commit 0da2da83a5
7 changed files with 74 additions and 25 deletions
@@ -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")
+20 -2
View File
@@ -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}")