mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-28 18:13:52 +08:00
update:检验LLM密钥
This commit is contained in:
@@ -2,6 +2,7 @@ from config.logger import setup_logging
|
|||||||
from http import HTTPStatus
|
from http import HTTPStatus
|
||||||
from dashscope import Application
|
from dashscope import Application
|
||||||
from core.providers.llm.base import LLMProviderBase
|
from core.providers.llm.base import LLMProviderBase
|
||||||
|
from core.utils.util import check_model_key
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -14,6 +15,7 @@ class LLMProvider(LLMProviderBase):
|
|||||||
self.base_url = config.get("base_url")
|
self.base_url = config.get("base_url")
|
||||||
self.is_No_prompt = config.get("is_no_prompt")
|
self.is_No_prompt = config.get("is_no_prompt")
|
||||||
self.memory_id = config.get("ali_memory_id")
|
self.memory_id = config.get("ali_memory_id")
|
||||||
|
check_model_key("AliBLLLM", self.api_key)
|
||||||
|
|
||||||
def response(self, session_id, dialogue):
|
def response(self, session_id, dialogue):
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -1,14 +1,17 @@
|
|||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
import requests
|
|
||||||
import json
|
import json
|
||||||
import re
|
|
||||||
from core.providers.llm.base import LLMProviderBase
|
from core.providers.llm.base import LLMProviderBase
|
||||||
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,
|
||||||
|
ChatEventType,
|
||||||
|
) # noqa
|
||||||
from core.providers.llm.system_prompt import get_system_prompt_for_function
|
from core.providers.llm.system_prompt import get_system_prompt_for_function
|
||||||
|
from core.utils.util import check_model_key
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -20,6 +23,7 @@ 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)
|
||||||
|
|
||||||
def response(self, session_id, dialogue):
|
def response(self, session_id, dialogue):
|
||||||
coze_api_token = self.personal_access_token
|
coze_api_token = self.personal_access_token
|
||||||
@@ -48,22 +52,22 @@ class LLMProvider(LLMProviderBase):
|
|||||||
print(event.message.content, end="", flush=True)
|
print(event.message.content, end="", flush=True)
|
||||||
yield event.message.content
|
yield event.message.content
|
||||||
|
|
||||||
def response_with_functions(self, session_id, dialogue, functions=None):
|
def response_with_functions(self, session_id, dialogue, functions=None):
|
||||||
if len(dialogue) == 2 and functions is not None and len(functions) > 0:
|
if len(dialogue) == 2 and functions is not None and len(functions) > 0:
|
||||||
# 第一次调用llm, 取最后一条用户消息,附加tool提示词
|
# 第一次调用llm, 取最后一条用户消息,附加tool提示词
|
||||||
last_msg = dialogue[-1]["content"]
|
last_msg = dialogue[-1]["content"]
|
||||||
function_str = json.dumps(functions, ensure_ascii=False)
|
function_str = json.dumps(functions, ensure_ascii=False)
|
||||||
modify_msg = get_system_prompt_for_function(function_str) + last_msg
|
modify_msg = get_system_prompt_for_function(function_str) + last_msg
|
||||||
dialogue[-1]["content"] = modify_msg
|
dialogue[-1]["content"] = modify_msg
|
||||||
|
|
||||||
# 如果最后一个是 role="tool",附加到user上
|
# 如果最后一个是 role="tool",附加到user上
|
||||||
if len(dialogue) > 1 and dialogue[-1]["role"] == "tool":
|
if len(dialogue) > 1 and dialogue[-1]["role"] == "tool":
|
||||||
assistant_msg = "\ntool call result: " + dialogue[-1]["content"] + "\n\n"
|
assistant_msg = "\ntool call result: " + dialogue[-1]["content"] + "\n\n"
|
||||||
while len(dialogue) > 1 :
|
while len(dialogue) > 1:
|
||||||
if dialogue[-1]["role"] == "user":
|
if dialogue[-1]["role"] == "user":
|
||||||
dialogue[-1]["content"] = assistant_msg + dialogue[-1]["content"]
|
dialogue[-1]["content"] = assistant_msg + dialogue[-1]["content"]
|
||||||
break
|
break
|
||||||
dialogue.pop()
|
dialogue.pop()
|
||||||
|
|
||||||
for token in self.response(session_id, dialogue):
|
for token in self.response(session_id, dialogue):
|
||||||
yield token, None
|
yield token, None
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ from config.logger import setup_logging
|
|||||||
import requests
|
import requests
|
||||||
from core.providers.llm.base import LLMProviderBase
|
from core.providers.llm.base import LLMProviderBase
|
||||||
from core.providers.llm.system_prompt import get_system_prompt_for_function
|
from core.providers.llm.system_prompt import get_system_prompt_for_function
|
||||||
|
from core.utils.util import check_model_key
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -14,6 +15,7 @@ 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)
|
||||||
|
|
||||||
def response(self, session_id, dialogue):
|
def response(self, session_id, dialogue):
|
||||||
try:
|
try:
|
||||||
@@ -60,7 +62,9 @@ class LLMProvider(LLMProviderBase):
|
|||||||
conversation_id # 更新映射
|
conversation_id # 更新映射
|
||||||
)
|
)
|
||||||
# 过滤 message_replace 事件,此事件会全量推一次
|
# 过滤 message_replace 事件,此事件会全量推一次
|
||||||
if event.get("event") != "message_replace" and event.get("answer"):
|
if event.get("event") != "message_replace" and event.get(
|
||||||
|
"answer"
|
||||||
|
):
|
||||||
yield event["answer"]
|
yield event["answer"]
|
||||||
elif self.mode == "workflows/run":
|
elif self.mode == "workflows/run":
|
||||||
for line in r.iter_lines():
|
for line in r.iter_lines():
|
||||||
@@ -76,29 +80,31 @@ class LLMProvider(LLMProviderBase):
|
|||||||
if line.startswith(b"data: "):
|
if line.startswith(b"data: "):
|
||||||
event = json.loads(line[6:])
|
event = json.loads(line[6:])
|
||||||
# 过滤 message_replace 事件,此事件会全量推一次
|
# 过滤 message_replace 事件,此事件会全量推一次
|
||||||
if event.get("event") != "message_replace" and event.get("answer"):
|
if event.get("event") != "message_replace" and event.get(
|
||||||
|
"answer"
|
||||||
|
):
|
||||||
yield event["answer"]
|
yield event["answer"]
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"Error in response generation: {e}")
|
logger.bind(tag=TAG).error(f"Error in response generation: {e}")
|
||||||
yield "【服务响应异常】"
|
yield "【服务响应异常】"
|
||||||
|
|
||||||
def response_with_functions(self, session_id, dialogue, functions=None):
|
def response_with_functions(self, session_id, dialogue, functions=None):
|
||||||
if len(dialogue) == 2 and functions is not None and len(functions) > 0:
|
if len(dialogue) == 2 and functions is not None and len(functions) > 0:
|
||||||
# 第一次调用llm, 取最后一条用户消息,附加tool提示词
|
# 第一次调用llm, 取最后一条用户消息,附加tool提示词
|
||||||
last_msg = dialogue[-1]["content"]
|
last_msg = dialogue[-1]["content"]
|
||||||
function_str = json.dumps(functions, ensure_ascii=False)
|
function_str = json.dumps(functions, ensure_ascii=False)
|
||||||
modify_msg = get_system_prompt_for_function(function_str) + last_msg
|
modify_msg = get_system_prompt_for_function(function_str) + last_msg
|
||||||
dialogue[-1]["content"] = modify_msg
|
dialogue[-1]["content"] = modify_msg
|
||||||
|
|
||||||
# 如果最后一个是 role="tool",附加到user上
|
# 如果最后一个是 role="tool",附加到user上
|
||||||
if len(dialogue) > 1 and dialogue[-1]["role"] == "tool":
|
if len(dialogue) > 1 and dialogue[-1]["role"] == "tool":
|
||||||
assistant_msg = "\ntool call result: " + dialogue[-1]["content"] + "\n\n"
|
assistant_msg = "\ntool call result: " + dialogue[-1]["content"] + "\n\n"
|
||||||
while len(dialogue) > 1 :
|
while len(dialogue) > 1:
|
||||||
if dialogue[-1]["role"] == "user":
|
if dialogue[-1]["role"] == "user":
|
||||||
dialogue[-1]["content"] = assistant_msg + dialogue[-1]["content"]
|
dialogue[-1]["content"] = assistant_msg + dialogue[-1]["content"]
|
||||||
break
|
break
|
||||||
dialogue.pop()
|
dialogue.pop()
|
||||||
|
|
||||||
for token in self.response(session_id, dialogue):
|
for token in self.response(session_id, dialogue):
|
||||||
yield token, None
|
yield token, None
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ import json
|
|||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
import requests
|
import requests
|
||||||
from core.providers.llm.base import LLMProviderBase
|
from core.providers.llm.base import LLMProviderBase
|
||||||
|
from core.utils.util import check_model_key
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -13,6 +14,7 @@ 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("LLM", self.api_key)
|
||||||
|
|
||||||
def response(self, session_id, dialogue):
|
def response(self, session_id, dialogue):
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -159,7 +159,6 @@ def check_model_key(modelType, modelKey):
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
"你还没配置" + modelType + "的密钥,请检查一下所使用的LLM是否配置了密钥"
|
"你还没配置" + modelType + "的密钥,请检查一下所使用的LLM是否配置了密钥"
|
||||||
)
|
)
|
||||||
return False
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
@@ -211,7 +210,7 @@ def check_ffmpeg_installed():
|
|||||||
def extract_json_from_string(input_string):
|
def extract_json_from_string(input_string):
|
||||||
"""提取字符串中的 JSON 部分"""
|
"""提取字符串中的 JSON 部分"""
|
||||||
pattern = r"(\{.*\})"
|
pattern = r"(\{.*\})"
|
||||||
match = re.search(pattern, input_string, re.DOTALL) #添加 re.DOTALL
|
match = re.search(pattern, input_string, re.DOTALL) # 添加 re.DOTALL
|
||||||
if match:
|
if match:
|
||||||
return match.group(1) # 返回提取的 JSON 字符串
|
return match.group(1) # 返回提取的 JSON 字符串
|
||||||
return None
|
return None
|
||||||
|
|||||||
Reference in New Issue
Block a user