mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-28 10:03:54 +08:00
update:优化空密钥的提示提示方式
This commit is contained in:
@@ -23,7 +23,9 @@ class LLMProvider(LLMProviderBase):
|
|||||||
self.bot_id = str(config.get("bot_id"))
|
self.bot_id = str(config.get("bot_id"))
|
||||||
self.user_id = str(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的映射
|
||||||
check_model_key("CozeLLM", self.personal_access_token)
|
model_key_msg = check_model_key("CozeLLM", self.personal_access_token)
|
||||||
|
if model_key_msg:
|
||||||
|
logger.bind(tag=TAG).error(model_key_msg)
|
||||||
|
|
||||||
def response(self, session_id, dialogue, **kwargs):
|
def response(self, session_id, dialogue, **kwargs):
|
||||||
coze_api_token = self.personal_access_token
|
coze_api_token = self.personal_access_token
|
||||||
|
|||||||
@@ -15,7 +15,9 @@ class LLMProvider(LLMProviderBase):
|
|||||||
self.mode = config.get("mode", "chat-messages")
|
self.mode = config.get("mode", "chat-messages")
|
||||||
self.base_url = config.get("base_url", "https://api.dify.ai/v1").rstrip("/")
|
self.base_url = config.get("base_url", "https://api.dify.ai/v1").rstrip("/")
|
||||||
self.session_conversation_map = {} # 存储session_id和conversation_id的映射
|
self.session_conversation_map = {} # 存储session_id和conversation_id的映射
|
||||||
check_model_key("DifyLLM", self.api_key)
|
model_key_msg = check_model_key("DifyLLM", self.api_key)
|
||||||
|
if model_key_msg:
|
||||||
|
logger.bind(tag=TAG).error(model_key_msg)
|
||||||
|
|
||||||
def response(self, session_id, dialogue, **kwargs):
|
def response(self, session_id, dialogue, **kwargs):
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -14,7 +14,9 @@ class LLMProvider(LLMProviderBase):
|
|||||||
self.base_url = config.get("base_url")
|
self.base_url = config.get("base_url")
|
||||||
self.detail = config.get("detail", False)
|
self.detail = config.get("detail", False)
|
||||||
self.variables = config.get("variables", {})
|
self.variables = config.get("variables", {})
|
||||||
check_model_key("FastGPTLLM", self.api_key)
|
model_key_msg = check_model_key("FastGPTLLM", self.api_key)
|
||||||
|
if model_key_msg:
|
||||||
|
logger.bind(tag=TAG).error(model_key_msg)
|
||||||
|
|
||||||
def response(self, session_id, dialogue, **kwargs):
|
def response(self, session_id, dialogue, **kwargs):
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -73,8 +73,9 @@ class LLMProvider(LLMProviderBase):
|
|||||||
http_proxy = cfg.get("http_proxy")
|
http_proxy = cfg.get("http_proxy")
|
||||||
https_proxy = cfg.get("https_proxy")
|
https_proxy = cfg.get("https_proxy")
|
||||||
|
|
||||||
if not check_model_key("LLM", self.api_key):
|
model_key_msg = check_model_key("LLM", self.api_key)
|
||||||
raise ValueError("无效的Gemini API Key,请检查是否配置正确")
|
if model_key_msg:
|
||||||
|
log.bind(tag=TAG).error(model_key_msg)
|
||||||
|
|
||||||
if http_proxy or https_proxy:
|
if http_proxy or https_proxy:
|
||||||
log.bind(tag=TAG).info(
|
log.bind(tag=TAG).info(
|
||||||
|
|||||||
@@ -21,20 +21,27 @@ class LLMProvider(LLMProviderBase):
|
|||||||
"max_tokens": (500, int),
|
"max_tokens": (500, int),
|
||||||
"temperature": (0.7, lambda x: round(float(x), 1)),
|
"temperature": (0.7, lambda x: round(float(x), 1)),
|
||||||
"top_p": (1.0, lambda x: round(float(x), 1)),
|
"top_p": (1.0, lambda x: round(float(x), 1)),
|
||||||
"frequency_penalty": (0, lambda x: round(float(x), 1))
|
"frequency_penalty": (0, lambda x: round(float(x), 1)),
|
||||||
}
|
}
|
||||||
|
|
||||||
for param, (default, converter) in param_defaults.items():
|
for param, (default, converter) in param_defaults.items():
|
||||||
value = config.get(param)
|
value = config.get(param)
|
||||||
try:
|
try:
|
||||||
setattr(self, param, converter(value) if value not in (None, "") else default)
|
setattr(
|
||||||
|
self,
|
||||||
|
param,
|
||||||
|
converter(value) if value not in (None, "") else default,
|
||||||
|
)
|
||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
setattr(self, param, default)
|
setattr(self, param, default)
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"意图识别参数初始化: {self.temperature}, {self.max_tokens}, {self.top_p}, {self.frequency_penalty}")
|
f"意图识别参数初始化: {self.temperature}, {self.max_tokens}, {self.top_p}, {self.frequency_penalty}"
|
||||||
|
)
|
||||||
|
|
||||||
check_model_key("LLM", self.api_key)
|
model_key_msg = check_model_key("LLM", self.api_key)
|
||||||
|
if model_key_msg:
|
||||||
|
logger.bind(tag=TAG).error(model_key_msg)
|
||||||
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
|
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
|
||||||
|
|
||||||
def response(self, session_id, dialogue, **kwargs):
|
def response(self, session_id, dialogue, **kwargs):
|
||||||
@@ -46,7 +53,9 @@ class LLMProvider(LLMProviderBase):
|
|||||||
max_tokens=kwargs.get("max_tokens", self.max_tokens),
|
max_tokens=kwargs.get("max_tokens", self.max_tokens),
|
||||||
temperature=kwargs.get("temperature", self.temperature),
|
temperature=kwargs.get("temperature", self.temperature),
|
||||||
top_p=kwargs.get("top_p", self.top_p),
|
top_p=kwargs.get("top_p", self.top_p),
|
||||||
frequency_penalty=kwargs.get("frequency_penalty", self.frequency_penalty),
|
frequency_penalty=kwargs.get(
|
||||||
|
"frequency_penalty", self.frequency_penalty
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
is_active = True
|
is_active = True
|
||||||
@@ -84,12 +93,14 @@ class LLMProvider(LLMProviderBase):
|
|||||||
for chunk in stream:
|
for chunk in stream:
|
||||||
# 检查是否存在有效的choice且content不为空
|
# 检查是否存在有效的choice且content不为空
|
||||||
if getattr(chunk, "choices", None):
|
if getattr(chunk, "choices", None):
|
||||||
yield chunk.choices[0].delta.content, chunk.choices[0].delta.tool_calls
|
yield chunk.choices[0].delta.content, chunk.choices[
|
||||||
|
0
|
||||||
|
].delta.tool_calls
|
||||||
# 存在 CompletionUsage 消息时,生成 Token 消耗 log
|
# 存在 CompletionUsage 消息时,生成 Token 消耗 log
|
||||||
elif isinstance(getattr(chunk, 'usage', None), CompletionUsage):
|
elif isinstance(getattr(chunk, "usage", None), CompletionUsage):
|
||||||
usage_info = getattr(chunk, 'usage', None)
|
usage_info = getattr(chunk, "usage", None)
|
||||||
logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).info(
|
||||||
f"Token 消耗:输入 {getattr(usage_info, 'prompt_tokens', '未知')},"
|
f"Token 消耗:输入 {getattr(usage_info, 'prompt_tokens', '未知')},"
|
||||||
f"输出 {getattr(usage_info, 'completion_tokens', '未知')},"
|
f"输出 {getattr(usage_info, 'completion_tokens', '未知')},"
|
||||||
f"共计 {getattr(usage_info, 'total_tokens', '未知')}"
|
f"共计 {getattr(usage_info, 'total_tokens', '未知')}"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -12,12 +12,14 @@ class MemoryProvider(MemoryProviderBase):
|
|||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
self.api_key = config.get("api_key", "")
|
self.api_key = config.get("api_key", "")
|
||||||
self.api_version = config.get("api_version", "v1.1")
|
self.api_version = config.get("api_version", "v1.1")
|
||||||
have_key = check_model_key("Mem0ai", self.api_key)
|
model_key_msg = check_model_key("Mem0ai", self.api_key)
|
||||||
if not have_key:
|
if model_key_msg:
|
||||||
|
logger.bind(tag=TAG).error(model_key_msg)
|
||||||
self.use_mem0 = False
|
self.use_mem0 = False
|
||||||
return
|
return
|
||||||
else:
|
else:
|
||||||
self.use_mem0 = True
|
self.use_mem0 = True
|
||||||
|
|
||||||
try:
|
try:
|
||||||
self.client = MemoryClient(api_key=self.api_key)
|
self.client = MemoryClient(api_key=self.api_key)
|
||||||
logger.bind(tag=TAG).info("成功连接到 Mem0ai 服务")
|
logger.bind(tag=TAG).info("成功连接到 Mem0ai 服务")
|
||||||
|
|||||||
@@ -37,7 +37,9 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.api_url = config.get("api_url")
|
self.api_url = config.get("api_url")
|
||||||
self.authorization = config.get("authorization")
|
self.authorization = config.get("authorization")
|
||||||
self.header = {"Authorization": f"{self.authorization}{self.access_token}"}
|
self.header = {"Authorization": f"{self.authorization}{self.access_token}"}
|
||||||
check_model_key("TTS", self.access_token)
|
model_key_msg = check_model_key("TTS", self.access_token)
|
||||||
|
if model_key_msg:
|
||||||
|
logger.bind(tag=TAG).error(model_key_msg)
|
||||||
|
|
||||||
async def text_to_speak(self, text, output_file):
|
async def text_to_speak(self, text, output_file):
|
||||||
request_json = {
|
request_json = {
|
||||||
|
|||||||
@@ -90,8 +90,9 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.format = config.get("response_format", "wav")
|
self.format = config.get("response_format", "wav")
|
||||||
self.audio_file_type = config.get("response_format", "wav")
|
self.audio_file_type = config.get("response_format", "wav")
|
||||||
self.api_key = config.get("api_key", "YOUR_API_KEY")
|
self.api_key = config.get("api_key", "YOUR_API_KEY")
|
||||||
have_key = check_model_key("FishSpeech TTS", self.api_key)
|
model_key_msg = check_model_key("FishSpeech TTS", self.api_key)
|
||||||
if not have_key:
|
if model_key_msg:
|
||||||
|
logger.bind(tag=TAG).error(model_key_msg)
|
||||||
return
|
return
|
||||||
self.normalize = str(config.get("normalize", True)).lower() in (
|
self.normalize = str(config.get("normalize", True)).lower() in (
|
||||||
"true",
|
"true",
|
||||||
|
|||||||
@@ -157,7 +157,9 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
||||||
sample_rate=16000, channels=1, frame_size_ms=60
|
sample_rate=16000, channels=1, frame_size_ms=60
|
||||||
)
|
)
|
||||||
check_model_key("TTS", self.access_token)
|
model_key_msg = check_model_key("TTS", self.access_token)
|
||||||
|
if model_key_msg:
|
||||||
|
logger.bind(tag=TAG).error(model_key_msg)
|
||||||
|
|
||||||
async def open_audio_channels(self, conn):
|
async def open_audio_channels(self, conn):
|
||||||
try:
|
try:
|
||||||
@@ -267,7 +269,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
await handleAbortMessage(self.conn)
|
await handleAbortMessage(self.conn)
|
||||||
logger.bind(tag=TAG).error(f"WebSocket连接不存在,终止发送文本")
|
logger.bind(tag=TAG).error(f"WebSocket连接不存在,终止发送文本")
|
||||||
return
|
return
|
||||||
|
|
||||||
# 过滤Markdown
|
# 过滤Markdown
|
||||||
filtered_text = MarkdownCleaner.clean_markdown(text)
|
filtered_text = MarkdownCleaner.clean_markdown(text)
|
||||||
|
|
||||||
|
|||||||
@@ -166,7 +166,9 @@ class TTSProvider(TTSProviderBase):
|
|||||||
) as resp:
|
) as resp:
|
||||||
|
|
||||||
if resp.status != 200:
|
if resp.status != 200:
|
||||||
logger.error(f"TTS请求失败: {resp.status}, {await resp.text()}")
|
logger.bind(tag=TAG).error(
|
||||||
|
f"TTS请求失败: {resp.status}, {await resp.text()}"
|
||||||
|
)
|
||||||
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -229,7 +231,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self._process_before_stop_play_files()
|
self._process_before_stop_play_files()
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"TTS请求异常: {e}")
|
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
|
||||||
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
||||||
|
|
||||||
def to_tts(self, text: str) -> list:
|
def to_tts(self, text: str) -> list:
|
||||||
@@ -263,7 +265,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.api_url, params=params, headers=headers, timeout=5
|
self.api_url, params=params, headers=headers, timeout=5
|
||||||
) as response:
|
) as response:
|
||||||
if response.status_code != 200:
|
if response.status_code != 200:
|
||||||
logger.error(
|
logger.bind(tag=TAG).error(
|
||||||
f"TTS请求失败: {response.status_code}, {response.text}"
|
f"TTS请求失败: {response.status_code}, {response.text}"
|
||||||
)
|
)
|
||||||
return []
|
return []
|
||||||
@@ -299,5 +301,5 @@ class TTSProvider(TTSProviderBase):
|
|||||||
return opus_datas
|
return opus_datas
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"TTS请求异常: {e}")
|
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
|
||||||
return []
|
return []
|
||||||
|
|||||||
@@ -25,7 +25,9 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.speed = float(speed) if speed else 1.0
|
self.speed = float(speed) if speed else 1.0
|
||||||
|
|
||||||
self.output_file = config.get("output_dir", "tmp/")
|
self.output_file = config.get("output_dir", "tmp/")
|
||||||
check_model_key("TTS", self.api_key)
|
model_key_msg = check_model_key("TTS", self.api_key)
|
||||||
|
if model_key_msg:
|
||||||
|
logger.bind(tag=TAG).error(model_key_msg)
|
||||||
|
|
||||||
async def text_to_speak(self, text, output_file):
|
async def text_to_speak(self, text, output_file):
|
||||||
headers = {
|
headers = {
|
||||||
|
|||||||
@@ -34,7 +34,9 @@ class VLLMProvider(VLLMProviderBase):
|
|||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
setattr(self, param, default)
|
setattr(self, param, default)
|
||||||
|
|
||||||
check_model_key("VLLM", self.api_key)
|
model_key_msg = check_model_key("VLLM", self.api_key)
|
||||||
|
if model_key_msg:
|
||||||
|
logger.bind(tag=TAG).error(model_key_msg)
|
||||||
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
|
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
|
||||||
|
|
||||||
def response(self, question, base64_image):
|
def response(self, question, base64_image):
|
||||||
|
|||||||
@@ -186,10 +186,8 @@ def remove_punctuation_and_length(text):
|
|||||||
|
|
||||||
def check_model_key(modelType, modelKey):
|
def check_model_key(modelType, modelKey):
|
||||||
if "你" in modelKey:
|
if "你" in modelKey:
|
||||||
raise ValueError(
|
return f"配置错误: {modelType} 的 API key 未设置,当前值为: {modelKey}"
|
||||||
"你还没配置" + modelType + "的密钥,请检查一下所使用的LLM是否配置了密钥"
|
return None
|
||||||
)
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|
||||||
def parse_string_to_list(value, separator=";"):
|
def parse_string_to_list(value, separator=";"):
|
||||||
@@ -785,7 +783,9 @@ def audio_bytes_to_data(audio_bytes, file_type, is_opus=True):
|
|||||||
return p3.decode_opus_from_bytes(audio_bytes)
|
return p3.decode_opus_from_bytes(audio_bytes)
|
||||||
else:
|
else:
|
||||||
# 其他格式用pydub
|
# 其他格式用pydub
|
||||||
audio = AudioSegment.from_file(BytesIO(audio_bytes), format=file_type, parameters=["-nostdin"])
|
audio = AudioSegment.from_file(
|
||||||
|
BytesIO(audio_bytes), format=file_type, parameters=["-nostdin"]
|
||||||
|
)
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
|
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
|
||||||
duration = len(audio) / 1000.0
|
duration = len(audio) / 1000.0
|
||||||
raw_data = audio.raw_data
|
raw_data = audio.raw_data
|
||||||
@@ -838,11 +838,11 @@ def opus_datas_to_wav_bytes(opus_datas, sample_rate=16000, channels=1):
|
|||||||
pcm = decoder.decode(opus_frame, frame_size)
|
pcm = decoder.decode(opus_frame, frame_size)
|
||||||
pcm_datas.append(pcm)
|
pcm_datas.append(pcm)
|
||||||
|
|
||||||
pcm_bytes = b''.join(pcm_datas)
|
pcm_bytes = b"".join(pcm_datas)
|
||||||
|
|
||||||
# 写入wav字节流
|
# 写入wav字节流
|
||||||
wav_buffer = BytesIO()
|
wav_buffer = BytesIO()
|
||||||
with wave.open(wav_buffer, 'wb') as wf:
|
with wave.open(wav_buffer, "wb") as wf:
|
||||||
wf.setnchannels(channels)
|
wf.setnchannels(channels)
|
||||||
wf.setsampwidth(2) # 16bit
|
wf.setsampwidth(2) # 16bit
|
||||||
wf.setframerate(sample_rate)
|
wf.setframerate(sample_rate)
|
||||||
|
|||||||
@@ -10,9 +10,17 @@ def append_devices_to_prompt(conn):
|
|||||||
funcs = conn.config["Intent"][conn.config["selected_module"]["Intent"]].get(
|
funcs = conn.config["Intent"][conn.config["selected_module"]["Intent"]].get(
|
||||||
"functions", []
|
"functions", []
|
||||||
)
|
)
|
||||||
|
|
||||||
|
config_source = (
|
||||||
|
"home_assistant"
|
||||||
|
if conn.config["plugins"].get("home_assistant")
|
||||||
|
else "hass_get_state"
|
||||||
|
)
|
||||||
|
|
||||||
if "hass_get_state" in funcs or "hass_set_state" in funcs:
|
if "hass_get_state" in funcs or "hass_set_state" in funcs:
|
||||||
prompt = "\n下面是我家智能设备列表(位置,设备名,entity_id),可以通过homeassistant控制\n"
|
prompt = "\n下面是我家智能设备列表(位置,设备名,entity_id),可以通过homeassistant控制\n"
|
||||||
devices = conn.config["plugins"]["home_assistant"].get("devices", [])
|
# TODO 分割被控设备
|
||||||
|
devices = conn.config["plugins"].get(config_source, {}).get("devices", [])
|
||||||
if len(devices) == 0:
|
if len(devices) == 0:
|
||||||
return
|
return
|
||||||
for device in devices:
|
for device in devices:
|
||||||
@@ -24,21 +32,26 @@ def append_devices_to_prompt(conn):
|
|||||||
|
|
||||||
def initialize_hass_handler(conn):
|
def initialize_hass_handler(conn):
|
||||||
ha_config = {}
|
ha_config = {}
|
||||||
if conn.load_function_plugin:
|
if not conn.load_function_plugin:
|
||||||
if conn.config["plugins"].get("home_assistant"):
|
return ha_config
|
||||||
ha_config["base_url"] = conn.config["plugins"]["home_assistant"].get(
|
|
||||||
"base_url"
|
# 确定配置来源
|
||||||
)
|
config_source = (
|
||||||
ha_config["api_key"] = conn.config["plugins"]["home_assistant"].get(
|
"home_assistant"
|
||||||
"api_key"
|
if conn.config["plugins"].get("home_assistant")
|
||||||
)
|
else "hass_get_state"
|
||||||
check_model_key("home_assistant", ha_config.get("api_key"))
|
)
|
||||||
elif conn.config["plugins"].get("hass_get_state"):
|
if not conn.config["plugins"].get(config_source):
|
||||||
ha_config["base_url"] = conn.config["plugins"]["hass_get_state"].get(
|
return ha_config
|
||||||
"base_url"
|
|
||||||
)
|
# 统一获取配置
|
||||||
ha_config["api_key"] = conn.config["plugins"]["hass_get_state"].get(
|
plugin_config = conn.config["plugins"][config_source]
|
||||||
"api_key"
|
ha_config["base_url"] = plugin_config.get("base_url")
|
||||||
)
|
ha_config["api_key"] = plugin_config.get("api_key")
|
||||||
check_model_key("home_assistant", ha_config.get("api_key"))
|
|
||||||
|
# 统一检查API密钥
|
||||||
|
model_key_msg = check_model_key("home_assistant", ha_config.get("api_key"))
|
||||||
|
if model_key_msg:
|
||||||
|
logger.bind(tag=TAG).error(model_key_msg)
|
||||||
|
|
||||||
return ha_config
|
return ha_config
|
||||||
|
|||||||
Reference in New Issue
Block a user