Merge pull request #1351 from xinnan-tech/py_memory_llm

update: 记忆模块使用独立LLM openai增加超参
This commit is contained in:
欣南科技
2025-05-27 15:29:22 +08:00
committed by GitHub
25 changed files with 324 additions and 155 deletions
+1 -1
View File
@@ -258,4 +258,4 @@
</plugin> </plugin>
</plugins> </plugins>
</build> </build>
</project> </project>
@@ -28,4 +28,6 @@ public class AgentChatHistoryReportDTO {
private String content; private String content;
@Schema(description = "base64编码的opus音频数据", example = "") @Schema(description = "base64编码的opus音频数据", example = "")
private String audioBase64; private String audioBase64;
@Schema(description = "上报时间,十位时间戳,空时默认使用当前时间", example = "1745657732")
private Long reportTime;
} }
@@ -27,4 +27,4 @@ public interface AgentChatAudioService extends IService<AgentChatAudioEntity> {
* @return 音频数据 * @return 音频数据
*/ */
byte[] getAudio(String audioId); byte[] getAudio(String audioId);
} }
@@ -47,7 +47,8 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic
public Boolean report(AgentChatHistoryReportDTO report) { public Boolean report(AgentChatHistoryReportDTO report) {
String macAddress = report.getMacAddress(); String macAddress = report.getMacAddress();
Byte chatType = report.getChatType(); Byte chatType = report.getChatType();
log.info("小智设备聊天上报请求: macAddress={}, type={}", macAddress, chatType); Long reportTimeMillis = null != report.getReportTime() ? report.getReportTime() * 1000 : System.currentTimeMillis();
log.info("小智设备聊天上报请求: macAddress={}, type={} reportTime={}", macAddress, chatType, reportTimeMillis);
// 根据设备MAC地址查询对应的默认智能体,判断是否需要上报 // 根据设备MAC地址查询对应的默认智能体,判断是否需要上报
AgentEntity agentEntity = agentService.getDefaultAgentByMacAddress(macAddress); AgentEntity agentEntity = agentService.getDefaultAgentByMacAddress(macAddress);
@@ -59,10 +60,10 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic
String agentId = agentEntity.getId(); String agentId = agentEntity.getId();
if (Objects.equals(chatHistoryConf, Constant.ChatHistoryConfEnum.RECORD_TEXT.getCode())) { if (Objects.equals(chatHistoryConf, Constant.ChatHistoryConfEnum.RECORD_TEXT.getCode())) {
saveChatText(report, agentId, macAddress, null); saveChatText(report, agentId, macAddress, null, reportTimeMillis);
} else if (Objects.equals(chatHistoryConf, Constant.ChatHistoryConfEnum.RECORD_TEXT_AUDIO.getCode())) { } else if (Objects.equals(chatHistoryConf, Constant.ChatHistoryConfEnum.RECORD_TEXT_AUDIO.getCode())) {
String audioId = saveChatAudio(report); String audioId = saveChatAudio(report);
saveChatText(report, agentId, macAddress, audioId); saveChatText(report, agentId, macAddress, audioId, reportTimeMillis);
} }
// 更新设备最后对话时间 // 更新设备最后对话时间
@@ -92,8 +93,7 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic
/** /**
* 组装上报数据 * 组装上报数据
*/ */
private void saveChatText(AgentChatHistoryReportDTO report, String agentId, String macAddress, String audioId) { private void saveChatText(AgentChatHistoryReportDTO report, String agentId, String macAddress, String audioId, Long reportTime) {
// 构建聊天记录实体 // 构建聊天记录实体
AgentChatHistoryEntity entity = AgentChatHistoryEntity.builder() AgentChatHistoryEntity entity = AgentChatHistoryEntity.builder()
.macAddress(macAddress) .macAddress(macAddress)
@@ -102,6 +102,8 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic
.chatType(report.getChatType()) .chatType(report.getChatType())
.content(report.getContent()) .content(report.getContent())
.audioId(audioId) .audioId(audioId)
.createdAt(new Date(reportTime))
// NOTE(haotian): 2025/5/26 updateAt可以不设置,重点是createAt,而且这样可以看到上报延迟
.build(); .build();
// 保存数据 // 保存数据
@@ -31,4 +31,4 @@ public class AgentChatAudioServiceImpl extends ServiceImpl<AiAgentChatAudioDao,
AgentChatAudioEntity entity = getById(audioId); AgentChatAudioEntity entity = getById(audioId);
return entity != null ? entity.getAudio() : null; return entity != null ? entity.getAudio() : null;
} }
} }
@@ -251,6 +251,7 @@ public class ConfigServiceImpl implements ConfigService {
String[] modelTypes = { "VAD", "ASR", "TTS", "Memory", "Intent", "LLM" }; String[] modelTypes = { "VAD", "ASR", "TTS", "Memory", "Intent", "LLM" };
String[] modelIds = { vadModelId, asrModelId, ttsModelId, memModelId, intentModelId, llmModelId }; String[] modelIds = { vadModelId, asrModelId, ttsModelId, memModelId, intentModelId, llmModelId };
String intentLLMModelId = null; String intentLLMModelId = null;
String memLocalShortLLMModelId = null;
for (int i = 0; i < modelIds.length; i++) { for (int i = 0; i < modelIds.length; i++) {
if (modelIds[i] == null) { if (modelIds[i] == null) {
@@ -269,7 +270,7 @@ public class ConfigServiceImpl implements ConfigService {
Map<String, Object> map = (Map<String, Object>) model.getConfigJson(); Map<String, Object> map = (Map<String, Object>) model.getConfigJson();
if ("intent_llm".equals(map.get("type"))) { if ("intent_llm".equals(map.get("type"))) {
intentLLMModelId = (String) map.get("llm"); intentLLMModelId = (String) map.get("llm");
if (intentLLMModelId != null && intentLLMModelId.equals(llmModelId)) { if (StringUtils.isNotBlank(intentLLMModelId) && intentLLMModelId.equals(llmModelId)) {
intentLLMModelId = null; intentLLMModelId = null;
} }
} }
@@ -281,10 +282,31 @@ public class ConfigServiceImpl implements ConfigService {
} }
} }
} }
if ("Memory".equals(modelTypes[i])) {
Map<String, Object> map = (Map<String, Object>) model.getConfigJson();
if ("mem_local_short".equals(map.get("type"))) {
memLocalShortLLMModelId = (String) map.get("llm");
if (StringUtils.isNotBlank(memLocalShortLLMModelId)
&& memLocalShortLLMModelId.equals(llmModelId)) {
memLocalShortLLMModelId = null;
}
}
}
// 如果是LLM类型,且intentLLMModelId不为空,则添加附加模型 // 如果是LLM类型,且intentLLMModelId不为空,则添加附加模型
if ("LLM".equals(modelTypes[i]) && intentLLMModelId != null) { if ("LLM".equals(modelTypes[i])) {
ModelConfigEntity intentLLM = modelConfigService.getModelById(intentLLMModelId, isCache); if (StringUtils.isNotBlank(intentLLMModelId)) {
typeConfig.put(intentLLM.getId(), intentLLM.getConfigJson()); if (!typeConfig.containsKey(intentLLMModelId)) {
ModelConfigEntity intentLLM = modelConfigService.getModelById(intentLLMModelId, isCache);
typeConfig.put(intentLLM.getId(), intentLLM.getConfigJson());
}
}
if (StringUtils.isNotBlank(memLocalShortLLMModelId)) {
if (!typeConfig.containsKey(memLocalShortLLMModelId)) {
ModelConfigEntity memLocalShortLLM = modelConfigService
.getModelById(memLocalShortLLMModelId, isCache);
typeConfig.put(memLocalShortLLM.getId(), memLocalShortLLM.getConfigJson());
}
}
} }
} }
result.put(modelTypes[i], typeConfig); result.put(modelTypes[i], typeConfig);
@@ -48,4 +48,4 @@ spring:
max-active: 8 # 连接池最大连接数(使用负值表示没有限制) max-active: 8 # 连接池最大连接数(使用负值表示没有限制)
max-idle: 8 # 连接池中的最大空闲连接 max-idle: 8 # 连接池中的最大空闲连接
min-idle: 0 # 连接池中的最小空闲连接 min-idle: 0 # 连接池中的最小空闲连接
shutdown-timeout: 100ms # 客户端优雅关闭的等待时间 shutdown-timeout: 100ms # 客户端优雅关闭的等待时间
@@ -0,0 +1,4 @@
-- 本地短期记忆配置可以设置独立的LLM
update `ai_model_provider` set fields = '[{"key":"llm","label":"LLM模型","type":"string"}]' where id = 'SYSTEM_Memory_mem_local_short';
update `ai_model_config` set config_json = '{\"type\": \"mem_local_short\", \"llm\": \"LLM_ChatGLMLLM\"}' where id = 'Memory_mem_local_short';
@@ -162,4 +162,11 @@ databaseChangeLog:
changes: changes:
- sqlFile: - sqlFile:
encoding: utf8 encoding: utf8
path: classpath:db/changelog/202505151451.sql path: classpath:db/changelog/202505151451.sql
- changeSet:
id: 202505271412
author: hrz
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202505271412.sql
+5 -1
View File
@@ -218,8 +218,12 @@ Memory:
# 不想使用记忆功能,可以使用nomem # 不想使用记忆功能,可以使用nomem
type: nomem type: nomem
mem_local_short: mem_local_short:
# 本地记忆功能,通过selected_module的llm总结,数据保存在本地,不会上传到服务器 # 本地记忆功能,通过selected_module的llm总结,数据保存在本地服务器,不会上传到外部服务器
type: mem_local_short type: mem_local_short
# 配备记忆存储独立的思考模型
# 如果这里不填,则会默认使用selected_module.LLM的模型作为意图识别的思考模型
# 如果你的不想使用selected_module.LLM记忆存储,这里最好使用独立的LLM作为意图识别,例如使用免费的ChatGLMLLM
llm: ChatGLMLLM
ASR: ASR:
FunASR: FunASR:
@@ -160,7 +160,7 @@ def save_mem_local_short(mac_address: str, short_momery: str) -> Optional[Dict]:
def report( def report(
mac_address: str, session_id: str, chat_type: int, content: str, audio mac_address: str, session_id: str, chat_type: int, content: str, audio, report_time
) -> Optional[Dict]: ) -> Optional[Dict]:
"""带熔断的业务方法示例""" """带熔断的业务方法示例"""
if not content or not ManageApiClient._instance: if not content or not ManageApiClient._instance:
@@ -174,6 +174,7 @@ def report(
"sessionId": session_id, "sessionId": session_id,
"chatType": chat_type, "chatType": chat_type,
"content": content, "content": content,
"reportTime": report_time,
"audioBase64": ( "audioBase64": (
base64.b64encode(audio).decode("utf-8") if audio else None base64.b64encode(audio).decode("utf-8") if audio else None
), ),
+50 -9
View File
@@ -89,12 +89,12 @@ class ConnectionHandler:
self.stop_event = threading.Event() self.stop_event = threading.Event()
self.tts_queue = queue.Queue() self.tts_queue = queue.Queue()
self.audio_play_queue = queue.Queue() self.audio_play_queue = queue.Queue()
self.executor = ThreadPoolExecutor(max_workers=10) self.executor = ThreadPoolExecutor(max_workers=5)
# 上报线程 # 添加上报线程
self.report_queue = queue.Queue() self.report_queue = queue.Queue()
self.report_thread = None self.report_thread = None
# TODO(haotian): 2025/5/12 可以通过修改此处,调节asr的上报和tts的上报 # 未来可以通过修改此处,调节asr的上报和tts的上报,目前默认都开启
self.report_asr_enable = self.read_config_from_api self.report_asr_enable = self.read_config_from_api
self.report_tts_enable = self.read_config_from_api self.report_tts_enable = self.read_config_from_api
@@ -454,6 +454,37 @@ class ConnectionHandler:
save_to_file=not self.read_config_from_api, save_to_file=not self.read_config_from_api,
) )
# 获取记忆总结配置
memory_config = self.config["Memory"]
memory_type = self.config["Memory"][self.config["selected_module"]["Memory"]][
"type"
]
# 如果使用 nomen,直接返回
if memory_type == "nomem":
return
# 使用 mem_local_short 模式
elif memory_type == "mem_local_short":
memory_llm_name = memory_config[self.config["selected_module"]["Memory"]][
"llm"
]
if memory_llm_name and memory_llm_name in self.config["LLM"]:
# 如果配置了专用LLM,则创建独立的LLM实例
from core.utils import llm as llm_utils
memory_llm_config = self.config["LLM"][memory_llm_name]
memory_llm_type = memory_llm_config.get("type", memory_llm_name)
memory_llm = llm_utils.create_instance(
memory_llm_type, memory_llm_config
)
self.logger.bind(tag=TAG).info(
f"为记忆总结创建了专用LLM: {memory_llm_name}, 类型: {memory_llm_type}"
)
self.memory.set_llm(memory_llm)
else:
# 否则使用主LLM
self.memory.set_llm(self.llm)
self.logger.bind(tag=TAG).info("使用主LLM作为意图识别模型")
def _initialize_intent(self): def _initialize_intent(self):
self.intent_type = self.config["Intent"][ self.intent_type = self.config["Intent"][
self.config["selected_module"]["Intent"] self.config["selected_module"]["Intent"]
@@ -889,16 +920,15 @@ class ConnectionHandler:
if item is None: # 检测毒丸对象 if item is None: # 检测毒丸对象
break break
type, text, audio_data = item type, text, audio_data, report_time = item
try: try:
# 执行上报(传入二进制数据) # 提交任务到线程池
report(self, type, text, audio_data) self.executor.submit(
self._process_report, type, text, audio_data, report_time
)
except Exception as e: except Exception as e:
self.logger.bind(tag=TAG).error(f"聊天记录上报线程异常: {e}") self.logger.bind(tag=TAG).error(f"聊天记录上报线程异常: {e}")
finally:
# 标记任务完成
self.report_queue.task_done()
except queue.Empty: except queue.Empty:
continue continue
except Exception as e: except Exception as e:
@@ -906,6 +936,17 @@ class ConnectionHandler:
self.logger.bind(tag=TAG).info("聊天记录上报线程已退出") self.logger.bind(tag=TAG).info("聊天记录上报线程已退出")
def _process_report(self, type, text, audio_data, report_time):
"""处理上报任务"""
try:
# 执行上报(传入二进制数据)
report(self, type, text, audio_data, report_time)
except Exception as e:
self.logger.bind(tag=TAG).error(f"上报处理异常: {e}")
finally:
# 标记任务完成
self.report_queue.task_done()
def speak_and_play(self, file_path, content, text_index=0): def speak_and_play(self, file_path, content, text_index=0):
if file_path is not None: if file_path is not None:
self.logger.bind(tag=TAG).info(f"无需tts转换: 从文件播放,{file_path}") self.logger.bind(tag=TAG).info(f"无需tts转换: 从文件播放,{file_path}")
@@ -8,6 +8,7 @@ TTS上报功能已集成到ConnectionHandler类中。
具体实现请参考core/connection.py中的相关代码。 具体实现请参考core/connection.py中的相关代码。
""" """
import time
import opuslib_next import opuslib_next
@@ -16,7 +17,7 @@ from config.manage_api_client import report as manage_report
TAG = __name__ TAG = __name__
def report(conn, type, text, opus_data): def report(conn, type, text, opus_data, report_time):
"""执行聊天记录上报操作 """执行聊天记录上报操作
Args: Args:
@@ -24,6 +25,7 @@ def report(conn, type, text, opus_data):
type: 上报类型,1为用户,2为智能体 type: 上报类型,1为用户,2为智能体
text: 合成文本 text: 合成文本
opus_data: opus音频数据 opus_data: opus音频数据
report_time: 上报时间
""" """
try: try:
if opus_data: if opus_data:
@@ -37,6 +39,7 @@ def report(conn, type, text, opus_data):
chat_type=type, chat_type=type,
content=text, content=text,
audio=audio_data, audio=audio_data,
report_time=report_time,
) )
except Exception as e: except Exception as e:
conn.logger.bind(tag=TAG).error(f"聊天记录上报失败: {e}") conn.logger.bind(tag=TAG).error(f"聊天记录上报失败: {e}")
@@ -104,12 +107,12 @@ def enqueue_tts_report(conn, text, opus_data):
try: try:
# 使用连接对象的队列,传入文本和二进制数据而非文件路径 # 使用连接对象的队列,传入文本和二进制数据而非文件路径
if conn.chat_history_conf == 2: if conn.chat_history_conf == 2:
conn.report_queue.put((2, text, opus_data)) conn.report_queue.put((2, text, opus_data, int(time.time())))
conn.logger.bind(tag=TAG).debug( conn.logger.bind(tag=TAG).debug(
f"TTS数据已加入上报队列: {conn.device_id}, 音频大小: {len(opus_data)} " f"TTS数据已加入上报队列: {conn.device_id}, 音频大小: {len(opus_data)} "
) )
else: else:
conn.report_queue.put((2, text, None)) conn.report_queue.put((2, text, None, int(time.time())))
conn.logger.bind(tag=TAG).debug( conn.logger.bind(tag=TAG).debug(
f"TTS数据已加入上报队列: {conn.device_id}, 不上报音频" f"TTS数据已加入上报队列: {conn.device_id}, 不上报音频"
) )
@@ -132,12 +135,12 @@ def enqueue_asr_report(conn, text, opus_data):
try: try:
# 使用连接对象的队列,传入文本和二进制数据而非文件路径 # 使用连接对象的队列,传入文本和二进制数据而非文件路径
if conn.chat_history_conf == 2: if conn.chat_history_conf == 2:
conn.report_queue.put((1, text, opus_data)) conn.report_queue.put((1, text, opus_data, int(time.time())))
conn.logger.bind(tag=TAG).debug( conn.logger.bind(tag=TAG).debug(
f"ASR数据已加入上报队列: {conn.device_id}, 音频大小: {len(opus_data)} " f"ASR数据已加入上报队列: {conn.device_id}, 音频大小: {len(opus_data)} "
) )
else: else:
conn.report_queue.put((1, text, None)) conn.report_queue.put((1, text, None, int(time.time())))
conn.logger.bind(tag=TAG).debug( conn.logger.bind(tag=TAG).debug(
f"ASR数据已加入上报队列: {conn.device_id}, 不上报音频" f"ASR数据已加入上报队列: {conn.device_id}, 不上报音频"
) )
@@ -10,7 +10,7 @@ class LLMProviderBase(ABC):
"""LLM response generator""" """LLM response generator"""
pass pass
def response_no_stream(self, system_prompt, user_prompt): def response_no_stream(self, system_prompt, user_prompt, **kwargs):
try: try:
# 构造对话格式 # 构造对话格式
dialogue = [ dialogue = [
@@ -18,7 +18,7 @@ class LLMProviderBase(ABC):
{"role": "user", "content": user_prompt} {"role": "user", "content": user_prompt}
] ]
result = "" result = ""
for part in self.response("", dialogue): for part in self.response("", dialogue, **kwargs):
result += part result += part
return result return result
@@ -30,7 +30,7 @@ class LLMProviderBase(ABC):
""" """
Default implementation for function calling (streaming) Default implementation for function calling (streaming)
This should be overridden by providers that support function calls This should be overridden by providers that support function calls
Returns: generator that yields either text tokens or a special function call token Returns: generator that yields either text tokens or a special function call token
""" """
# For providers that don't support functions, just return regular response # For providers that don't support functions, just return regular response
@@ -25,7 +25,7 @@ class LLMProvider(LLMProviderBase):
self.session_conversation_map = {} # 存储session_id和conversation_id的映射 self.session_conversation_map = {} # 存储session_id和conversation_id的映射
check_model_key("CozeLLM", self.personal_access_token) check_model_key("CozeLLM", self.personal_access_token)
def response(self, session_id, dialogue): def response(self, session_id, dialogue, **kwargs):
coze_api_token = self.personal_access_token coze_api_token = self.personal_access_token
coze_api_base = COZE_CN_BASE_URL coze_api_base = COZE_CN_BASE_URL
@@ -17,7 +17,7 @@ class LLMProvider(LLMProviderBase):
self.session_conversation_map = {} # 存储session_id和conversation_id的映射 self.session_conversation_map = {} # 存储session_id和conversation_id的映射
check_model_key("DifyLLM", self.api_key) check_model_key("DifyLLM", self.api_key)
def response(self, session_id, dialogue): def response(self, session_id, dialogue, **kwargs):
try: try:
# 取最后一条用户消息 # 取最后一条用户消息
last_msg = next(m for m in reversed(dialogue) if m["role"] == "user") last_msg = next(m for m in reversed(dialogue) if m["role"] == "user")
@@ -16,7 +16,7 @@ class LLMProvider(LLMProviderBase):
self.variables = config.get("variables", {}) self.variables = config.get("variables", {})
check_model_key("FastGPTLLM", self.api_key) check_model_key("FastGPTLLM", self.api_key)
def response(self, session_id, dialogue): def response(self, session_id, dialogue, **kwargs):
try: try:
# 取最后一条用户消息 # 取最后一条用户消息
last_msg = next(m for m in reversed(dialogue) if m["role"] == "user") last_msg = next(m for m in reversed(dialogue) if m["role"] == "user")
@@ -112,7 +112,7 @@ class LLMProvider(LLMProviderBase):
] ]
# Gemini文档提到,无需维护session-id,直接用dialogue拼接而成 # Gemini文档提到,无需维护session-id,直接用dialogue拼接而成
def response(self, session_id, dialogue): def response(self, session_id, dialogue, **kwargs):
yield from self._generate(dialogue, None) yield from self._generate(dialogue, None)
def response_with_functions(self, session_id, dialogue, functions=None): def response_with_functions(self, session_id, dialogue, functions=None):
@@ -14,7 +14,7 @@ class LLMProvider(LLMProviderBase):
self.base_url = config.get("base_url", config.get("url")) # 默认使用 base_url self.base_url = config.get("base_url", config.get("url")) # 默认使用 base_url
self.api_url = f"{self.base_url}/api/conversation/process" # 拼接完整的 API URL self.api_url = f"{self.base_url}/api/conversation/process" # 拼接完整的 API URL
def response(self, session_id, dialogue): def response(self, session_id, dialogue, **kwargs):
try: try:
# home assistant语音助手自带意图,无需使用xiaozhi ai自带的,只需要把用户说的话传递给home assistant即可 # home assistant语音助手自带意图,无需使用xiaozhi ai自带的,只需要把用户说的话传递给home assistant即可
@@ -18,13 +18,13 @@ class LLMProvider(LLMProviderBase):
self.client = OpenAI( self.client = OpenAI(
base_url=self.base_url, base_url=self.base_url,
api_key="ollama" # Ollama doesn't need an API key but OpenAI client requires one api_key="ollama", # Ollama doesn't need an API key but OpenAI client requires one
) )
# 检查是否是qwen3模型 # 检查是否是qwen3模型
self.is_qwen3 = self.model_name and self.model_name.lower().startswith("qwen3") self.is_qwen3 = self.model_name and self.model_name.lower().startswith("qwen3")
def response(self, session_id, dialogue): def response(self, session_id, dialogue, **kwargs):
try: try:
# 如果是qwen3模型,在用户最后一条消息中添加/no_think指令 # 如果是qwen3模型,在用户最后一条消息中添加/no_think指令
if self.is_qwen3: if self.is_qwen3:
@@ -35,7 +35,9 @@ class LLMProvider(LLMProviderBase):
for i in range(len(dialogue_copy) - 1, -1, -1): for i in range(len(dialogue_copy) - 1, -1, -1):
if dialogue_copy[i]["role"] == "user": if dialogue_copy[i]["role"] == "user":
# 在用户消息前添加/no_think指令 # 在用户消息前添加/no_think指令
dialogue_copy[i]["content"] = "/no_think " + dialogue_copy[i]["content"] dialogue_copy[i]["content"] = (
"/no_think " + dialogue_copy[i]["content"]
)
logger.bind(tag=TAG).debug(f"为qwen3模型添加/no_think指令") logger.bind(tag=TAG).debug(f"为qwen3模型添加/no_think指令")
break break
@@ -43,9 +45,7 @@ class LLMProvider(LLMProviderBase):
dialogue = dialogue_copy dialogue = dialogue_copy
responses = self.client.chat.completions.create( responses = self.client.chat.completions.create(
model=self.model_name, model=self.model_name, messages=dialogue, stream=True
messages=dialogue,
stream=True
) )
is_active = True is_active = True
# 用于处理跨chunk的标签 # 用于处理跨chunk的标签
@@ -53,29 +53,33 @@ class LLMProvider(LLMProviderBase):
for chunk in responses: for chunk in responses:
try: try:
delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None delta = (
content = delta.content if hasattr(delta, 'content') else '' chunk.choices[0].delta
if getattr(chunk, "choices", None)
else None
)
content = delta.content if hasattr(delta, "content") else ""
if content: if content:
# 将内容添加到缓冲区 # 将内容添加到缓冲区
buffer += content buffer += content
# 处理缓冲区中的标签 # 处理缓冲区中的标签
while '<think>' in buffer and '</think>' in buffer: while "<think>" in buffer and "</think>" in buffer:
# 找到完整的<think></think>标签并移除 # 找到完整的<think></think>标签并移除
pre = buffer.split('<think>', 1)[0] pre = buffer.split("<think>", 1)[0]
post = buffer.split('</think>', 1)[1] post = buffer.split("</think>", 1)[1]
buffer = pre + post buffer = pre + post
# 处理只有开始标签的情况 # 处理只有开始标签的情况
if '<think>' in buffer: if "<think>" in buffer:
is_active = False is_active = False
buffer = buffer.split('<think>', 1)[0] buffer = buffer.split("<think>", 1)[0]
# 处理只有结束标签的情况 # 处理只有结束标签的情况
if '</think>' in buffer: if "</think>" in buffer:
is_active = True is_active = True
buffer = buffer.split('</think>', 1)[1] buffer = buffer.split("</think>", 1)[1]
# 如果当前处于活动状态且缓冲区有内容,则输出 # 如果当前处于活动状态且缓冲区有内容,则输出
if is_active and buffer: if is_active and buffer:
@@ -100,7 +104,9 @@ class LLMProvider(LLMProviderBase):
for i in range(len(dialogue_copy) - 1, -1, -1): for i in range(len(dialogue_copy) - 1, -1, -1):
if dialogue_copy[i]["role"] == "user": if dialogue_copy[i]["role"] == "user":
# 在用户消息前添加/no_think指令 # 在用户消息前添加/no_think指令
dialogue_copy[i]["content"] = "/no_think " + dialogue_copy[i]["content"] dialogue_copy[i]["content"] = (
"/no_think " + dialogue_copy[i]["content"]
)
logger.bind(tag=TAG).debug(f"为qwen3模型添加/no_think指令") logger.bind(tag=TAG).debug(f"为qwen3模型添加/no_think指令")
break break
@@ -119,9 +125,15 @@ class LLMProvider(LLMProviderBase):
for chunk in stream: for chunk in stream:
try: try:
delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None delta = (
content = delta.content if hasattr(delta, 'content') else None chunk.choices[0].delta
tool_calls = delta.tool_calls if hasattr(delta, 'tool_calls') else None if getattr(chunk, "choices", None)
else None
)
content = delta.content if hasattr(delta, "content") else None
tool_calls = (
delta.tool_calls if hasattr(delta, "tool_calls") else None
)
# 如果是工具调用,直接传递 # 如果是工具调用,直接传递
if tool_calls: if tool_calls:
@@ -134,21 +146,21 @@ class LLMProvider(LLMProviderBase):
buffer += content buffer += content
# 处理缓冲区中的标签 # 处理缓冲区中的标签
while '<think>' in buffer and '</think>' in buffer: while "<think>" in buffer and "</think>" in buffer:
# 找到完整的<think></think>标签并移除 # 找到完整的<think></think>标签并移除
pre = buffer.split('<think>', 1)[0] pre = buffer.split("<think>", 1)[0]
post = buffer.split('</think>', 1)[1] post = buffer.split("</think>", 1)[1]
buffer = pre + post buffer = pre + post
# 处理只有开始标签的情况 # 处理只有开始标签的情况
if '<think>' in buffer: if "<think>" in buffer:
is_active = False is_active = False
buffer = buffer.split('<think>', 1)[0] buffer = buffer.split("<think>", 1)[0]
# 处理只有结束标签的情况 # 处理只有结束标签的情况
if '</think>' in buffer: if "</think>" in buffer:
is_active = True is_active = True
buffer = buffer.split('</think>', 1)[1] buffer = buffer.split("</think>", 1)[1]
# 如果当前处于活动状态且缓冲区有内容,则输出 # 如果当前处于活动状态且缓冲区有内容,则输出
if is_active and buffer: if is_active and buffer:
@@ -16,26 +16,37 @@ class LLMProvider(LLMProviderBase):
self.base_url = config.get("base_url") self.base_url = config.get("base_url")
else: else:
self.base_url = config.get("url") self.base_url = config.get("url")
max_tokens = config.get("max_tokens")
if max_tokens is None or max_tokens == "":
max_tokens = 500
try: param_defaults = {
max_tokens = int(max_tokens) "max_tokens": (500, int),
except (ValueError, TypeError): "temperature": (0.7, lambda x: round(float(x), 1)),
max_tokens = 500 "top_p": (1.0, lambda x: round(float(x), 1)),
self.max_tokens = max_tokens "frequency_penalty": (0, lambda x: round(float(x), 1))
}
for param, (default, converter) in param_defaults.items():
value = config.get(param)
try:
setattr(self, param, converter(value) if value not in (None, "") else default)
except (ValueError, TypeError):
setattr(self, param, default)
logger.debug(
f"意图识别参数初始化: {self.temperature}, {self.max_tokens}, {self.top_p}, {self.frequency_penalty}")
check_model_key("LLM", self.api_key) check_model_key("LLM", self.api_key)
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): def response(self, session_id, dialogue, **kwargs):
try: try:
responses = self.client.chat.completions.create( responses = self.client.chat.completions.create(
model=self.model_name, model=self.model_name,
messages=dialogue, messages=dialogue,
stream=True, stream=True,
max_tokens=self.max_tokens, max_tokens=kwargs.get("max_tokens", self.max_tokens),
temperature=kwargs.get("temperature", self.temperature),
top_p=kwargs.get("top_p", self.top_p),
frequency_penalty=kwargs.get("frequency_penalty", self.frequency_penalty),
) )
is_active = True is_active = True
@@ -15,39 +15,45 @@ class LLMProvider(LLMProviderBase):
# 如果没有v1,增加v1 # 如果没有v1,增加v1
if not self.base_url.endswith("/v1"): if not self.base_url.endswith("/v1"):
self.base_url = f"{self.base_url}/v1" self.base_url = f"{self.base_url}/v1"
logger.bind(tag=TAG).info(f"Initializing Xinference LLM provider with model: {self.model_name}, base_url: {self.base_url}") logger.bind(tag=TAG).info(
f"Initializing Xinference LLM provider with model: {self.model_name}, base_url: {self.base_url}"
)
try: try:
self.client = OpenAI( self.client = OpenAI(
base_url=self.base_url, base_url=self.base_url,
api_key="xinference" # Xinference has a similar setup to Ollama where it doesn't need an actual key api_key="xinference", # Xinference has a similar setup to Ollama where it doesn't need an actual key
) )
logger.bind(tag=TAG).info("Xinference client initialized successfully") logger.bind(tag=TAG).info("Xinference client initialized successfully")
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"Error initializing Xinference client: {e}") logger.bind(tag=TAG).error(f"Error initializing Xinference client: {e}")
raise raise
def response(self, session_id, dialogue): def response(self, session_id, dialogue, **kwargs):
try: try:
logger.bind(tag=TAG).debug(f"Sending request to Xinference with model: {self.model_name}, dialogue length: {len(dialogue)}") logger.bind(tag=TAG).debug(
responses = self.client.chat.completions.create( f"Sending request to Xinference with model: {self.model_name}, dialogue length: {len(dialogue)}"
model=self.model_name,
messages=dialogue,
stream=True
) )
is_active=True responses = self.client.chat.completions.create(
model=self.model_name, messages=dialogue, stream=True
)
is_active = True
for chunk in responses: for chunk in responses:
try: try:
delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None delta = (
content = delta.content if hasattr(delta, 'content') else '' chunk.choices[0].delta
if getattr(chunk, "choices", None)
else None
)
content = delta.content if hasattr(delta, "content") else ""
if content: if content:
if '<think>' in content: if "<think>" in content:
is_active = False is_active = False
content = content.split('<think>')[0] content = content.split("<think>")[0]
if '</think>' in content: if "</think>" in content:
is_active = True is_active = True
content = content.split('</think>')[-1] content = content.split("</think>")[-1]
if is_active: if is_active:
yield content yield content
except Exception as e: except Exception as e:
@@ -59,10 +65,14 @@ class LLMProvider(LLMProviderBase):
def response_with_functions(self, session_id, dialogue, functions=None): def response_with_functions(self, session_id, dialogue, functions=None):
try: try:
logger.bind(tag=TAG).debug(f"Sending function call request to Xinference with model: {self.model_name}, dialogue length: {len(dialogue)}") logger.bind(tag=TAG).debug(
f"Sending function call request to Xinference with model: {self.model_name}, dialogue length: {len(dialogue)}"
)
if functions: if functions:
logger.bind(tag=TAG).debug(f"Function calls enabled with: {[f.get('function', {}).get('name') for f in functions]}") logger.bind(tag=TAG).debug(
f"Function calls enabled with: {[f.get('function', {}).get('name') for f in functions]}"
)
stream = self.client.chat.completions.create( stream = self.client.chat.completions.create(
model=self.model_name, model=self.model_name,
messages=dialogue, messages=dialogue,
@@ -74,7 +84,7 @@ class LLMProvider(LLMProviderBase):
delta = chunk.choices[0].delta delta = chunk.choices[0].delta
content = delta.content content = delta.content
tool_calls = delta.tool_calls tool_calls = delta.tool_calls
if content: if content:
yield content, tool_calls yield content, tool_calls
elif tool_calls: elif tool_calls:
@@ -82,4 +92,7 @@ class LLMProvider(LLMProviderBase):
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"Error in Xinference function call: {e}") logger.bind(tag=TAG).error(f"Error in Xinference function call: {e}")
yield {"type": "content", "content": f"【Xinference服务响应异常: {str(e)}"} yield {
"type": "content",
"content": f"【Xinference服务响应异常: {str(e)}",
}
@@ -9,7 +9,13 @@ class MemoryProviderBase(ABC):
def __init__(self, config): def __init__(self, config):
self.config = config self.config = config
self.role_id = None self.role_id = None
self.llm = None
def set_llm(self, llm):
self.llm = llm
# 获取模型名称和类型信息
model_name = getattr(llm, "model_name", str(llm.__class__.__name__))
# 记录更详细的日志
logger.bind(tag=TAG).info(f"记忆总结设置LLM: {model_name}")
@abstractmethod @abstractmethod
async def save_memory(self, msgs): async def save_memory(self, msgs):
@@ -107,7 +107,7 @@ TAG = __name__
class MemoryProvider(MemoryProviderBase): class MemoryProvider(MemoryProviderBase):
def __init__(self, config, summary_memory): def __init__(self, config, summary_memory):
super().__init__(config) super().__init__(config)
self.short_momery = "" self.short_memory = ""
self.save_to_file = True self.save_to_file = True
self.memory_path = get_project_dir() + "data/.memory.yaml" self.memory_path = get_project_dir() + "data/.memory.yaml"
self.load_memory(summary_memory) self.load_memory(summary_memory)
@@ -122,7 +122,7 @@ class MemoryProvider(MemoryProviderBase):
def load_memory(self, summary_memory): def load_memory(self, summary_memory):
# api获取到总结记忆后直接返回 # api获取到总结记忆后直接返回
if summary_memory or not self.save_to_file: if summary_memory or not self.save_to_file:
self.short_momery = summary_memory self.short_memory = summary_memory
return return
all_memory = {} all_memory = {}
@@ -130,18 +130,21 @@ class MemoryProvider(MemoryProviderBase):
with open(self.memory_path, "r", encoding="utf-8") as f: with open(self.memory_path, "r", encoding="utf-8") as f:
all_memory = yaml.safe_load(f) or {} all_memory = yaml.safe_load(f) or {}
if self.role_id in all_memory: if self.role_id in all_memory:
self.short_momery = all_memory[self.role_id] self.short_memory = all_memory[self.role_id]
def save_memory_to_file(self): def save_memory_to_file(self):
all_memory = {} all_memory = {}
if os.path.exists(self.memory_path): if os.path.exists(self.memory_path):
with open(self.memory_path, "r", encoding="utf-8") as f: with open(self.memory_path, "r", encoding="utf-8") as f:
all_memory = yaml.safe_load(f) or {} all_memory = yaml.safe_load(f) or {}
all_memory[self.role_id] = self.short_momery all_memory[self.role_id] = self.short_memory
with open(self.memory_path, "w", encoding="utf-8") as f: with open(self.memory_path, "w", encoding="utf-8") as f:
yaml.dump(all_memory, f, allow_unicode=True) yaml.dump(all_memory, f, allow_unicode=True)
async def save_memory(self, msgs): async def save_memory(self, msgs):
# 打印使用的模型信息
model_info = getattr(self.llm, "model_name", str(self.llm.__class__.__name__))
logger.bind(tag=TAG).debug(f"使用记忆保存模型: {model_info}")
if self.llm is None: if self.llm is None:
logger.bind(tag=TAG).error("LLM is not set for memory provider") logger.bind(tag=TAG).error("LLM is not set for memory provider")
return None return None
@@ -155,31 +158,39 @@ class MemoryProvider(MemoryProviderBase):
msgStr += f"User: {msg.content}\n" msgStr += f"User: {msg.content}\n"
elif msg.role == "assistant": elif msg.role == "assistant":
msgStr += f"Assistant: {msg.content}\n" msgStr += f"Assistant: {msg.content}\n"
if self.short_momery and len(self.short_momery) > 0: if self.short_memory and len(self.short_memory) > 0:
msgStr += "历史记忆:\n" msgStr += "历史记忆:\n"
msgStr += self.short_momery msgStr += self.short_memory
# 当前时间 # 当前时间
time_str = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime()) time_str = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime())
msgStr += f"当前时间:{time_str}" msgStr += f"当前时间:{time_str}"
if self.save_to_file: if self.save_to_file:
result = self.llm.response_no_stream(short_term_memory_prompt, msgStr) result = self.llm.response_no_stream(
short_term_memory_prompt,
msgStr,
max_tokens=2000,
temperature=0.2,
)
json_str = extract_json_data(result) json_str = extract_json_data(result)
try: try:
json.loads(json_str) # 检查json格式是否正确 json.loads(json_str) # 检查json格式是否正确
self.short_momery = json_str self.short_memory = json_str
self.save_memory_to_file() self.save_memory_to_file()
except Exception as e: except Exception as e:
print("Error:", e) print("Error:", e)
else: else:
result = self.llm.response_no_stream( result = self.llm.response_no_stream(
short_term_memory_prompt_only_content, msgStr short_term_memory_prompt_only_content,
msgStr,
max_tokens=2000,
temperature=0.2,
) )
save_mem_local_short(self.role_id, result) save_mem_local_short(self.role_id, result)
logger.bind(tag=TAG).info(f"Save memory successful - Role: {self.role_id}") logger.bind(tag=TAG).info(f"Save memory successful - Role: {self.role_id}")
return self.short_momery return self.short_memory
async def query_memory(self, query: str) -> str: async def query_memory(self, query: str) -> str:
return self.short_momery return self.short_memory
@@ -13,56 +13,71 @@ import urllib.parse
import time import time
import uuid import uuid
from urllib import parse from urllib import parse
class AccessToken: class AccessToken:
@staticmethod @staticmethod
def _encode_text(text): def _encode_text(text):
encoded_text = parse.quote_plus(text) encoded_text = parse.quote_plus(text)
return encoded_text.replace('+', '%20').replace('*', '%2A').replace('%7E', '~') return encoded_text.replace("+", "%20").replace("*", "%2A").replace("%7E", "~")
@staticmethod @staticmethod
def _encode_dict(dic): def _encode_dict(dic):
keys = dic.keys() keys = dic.keys()
dic_sorted = [(key, dic[key]) for key in sorted(keys)] dic_sorted = [(key, dic[key]) for key in sorted(keys)]
encoded_text = parse.urlencode(dic_sorted) encoded_text = parse.urlencode(dic_sorted)
return encoded_text.replace('+', '%20').replace('*', '%2A').replace('%7E', '~') return encoded_text.replace("+", "%20").replace("*", "%2A").replace("%7E", "~")
@staticmethod @staticmethod
def create_token(access_key_id, access_key_secret): def create_token(access_key_id, access_key_secret):
parameters = {'AccessKeyId': access_key_id, parameters = {
'Action': 'CreateToken', "AccessKeyId": access_key_id,
'Format': 'JSON', "Action": "CreateToken",
'RegionId': 'cn-shanghai', "Format": "JSON",
'SignatureMethod': 'HMAC-SHA1', "RegionId": "cn-shanghai",
'SignatureNonce': str(uuid.uuid1()), "SignatureMethod": "HMAC-SHA1",
'SignatureVersion': '1.0', "SignatureNonce": str(uuid.uuid1()),
'Timestamp': time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), "SignatureVersion": "1.0",
'Version': '2019-02-28'} "Timestamp": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
"Version": "2019-02-28",
}
# 构造规范化的请求字符串 # 构造规范化的请求字符串
query_string = AccessToken._encode_dict(parameters) query_string = AccessToken._encode_dict(parameters)
# print('规范化的请求字符串: %s' % query_string) # print('规范化的请求字符串: %s' % query_string)
# 构造待签名字符串 # 构造待签名字符串
string_to_sign = 'GET' + '&' + AccessToken._encode_text('/') + '&' + AccessToken._encode_text(query_string) string_to_sign = (
"GET"
+ "&"
+ AccessToken._encode_text("/")
+ "&"
+ AccessToken._encode_text(query_string)
)
# print('待签名的字符串: %s' % string_to_sign) # print('待签名的字符串: %s' % string_to_sign)
# 计算签名 # 计算签名
secreted_string = hmac.new(bytes(access_key_secret + '&', encoding='utf-8'), secreted_string = hmac.new(
bytes(string_to_sign, encoding='utf-8'), bytes(access_key_secret + "&", encoding="utf-8"),
hashlib.sha1).digest() bytes(string_to_sign, encoding="utf-8"),
hashlib.sha1,
).digest()
signature = base64.b64encode(secreted_string) signature = base64.b64encode(secreted_string)
# print('签名: %s' % signature) # print('签名: %s' % signature)
# 进行URL编码 # 进行URL编码
signature = AccessToken._encode_text(signature) signature = AccessToken._encode_text(signature)
# print('URL编码后的签名: %s' % signature) # print('URL编码后的签名: %s' % signature)
# 调用服务 # 调用服务
full_url = 'http://nls-meta.cn-shanghai.aliyuncs.com/?Signature=%s&%s' % (signature, query_string) full_url = "http://nls-meta.cn-shanghai.aliyuncs.com/?Signature=%s&%s" % (
signature,
query_string,
)
# print('url: %s' % full_url) # print('url: %s' % full_url)
# 提交HTTP GET请求 # 提交HTTP GET请求
response = requests.get(full_url) response = requests.get(full_url)
if response.ok: if response.ok:
root_obj = response.json() root_obj = response.json()
key = 'Token' key = "Token"
if key in root_obj: if key in root_obj:
token = root_obj[key]['Id'] token = root_obj[key]["Id"]
expire_time = root_obj[key]['ExpireTime'] expire_time = root_obj[key]["ExpireTime"]
return token, expire_time return token, expire_time
# print(response.text) # print(response.text)
return None, None return None, None
@@ -70,26 +85,36 @@ class AccessToken:
class TTSProvider(TTSProviderBase): class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file): def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file) super().__init__(config, delete_audio_file)
# 新增空值判断逻辑 # 新增空值判断逻辑
self.access_key_id = config.get("access_key_id") self.access_key_id = config.get("access_key_id")
self.access_key_secret = config.get("access_key_secret") self.access_key_secret = config.get("access_key_secret")
self.appkey = config.get("appkey") self.appkey = config.get("appkey")
self.format = config.get("format", "wav") self.format = config.get("format", "wav")
self.sample_rate = config.get("sample_rate", 16000)
self.voice = config.get("voice", "xiaoyun") sample_rate = config.get("sample_rate", "16000")
self.volume = config.get("volume", 50) self.sample_rate = int(sample_rate) if sample_rate else 16000
self.speech_rate = config.get("speech_rate", 0)
self.pitch_rate = config.get("pitch_rate", 0) if config.get("private_voice"):
self.voice = config.get("private_voice")
else:
self.voice = config.get("voice", "xiaoyun")
volume = config.get("volume", "50")
self.volume = int(volume) if volume else 50
speech_rate = config.get("speech_rate", "0")
self.speech_rate = int(speech_rate) if speech_rate else 0
pitch_rate = config.get("pitch_rate", "0")
self.pitch_rate = int(pitch_rate) if pitch_rate else 0
self.host = config.get("host", "nls-gateway-cn-shanghai.aliyuncs.com") self.host = config.get("host", "nls-gateway-cn-shanghai.aliyuncs.com")
self.api_url = f"https://{self.host}/stream/v1/tts" self.api_url = f"https://{self.host}/stream/v1/tts"
self.header = { self.header = {"Content-Type": "application/json"}
"Content-Type": "application/json"
}
if self.access_key_id and self.access_key_secret: if self.access_key_id and self.access_key_secret:
# 使用密钥对生成临时token # 使用密钥对生成临时token
@@ -99,35 +124,30 @@ class TTSProvider(TTSProviderBase):
self.token = config.get("token") self.token = config.get("token")
self.expire_time = None self.expire_time = None
def _refresh_token(self): def _refresh_token(self):
"""刷新Token并记录过期时间""" """刷新Token并记录过期时间"""
if self.access_key_id and self.access_key_secret: if self.access_key_id and self.access_key_secret:
self.token, expire_time_str = AccessToken.create_token( self.token, expire_time_str = AccessToken.create_token(
self.access_key_id, self.access_key_id, self.access_key_secret
self.access_key_secret
) )
if not expire_time_str: if not expire_time_str:
raise ValueError("无法获取有效的Token过期时间") raise ValueError("无法获取有效的Token过期时间")
try: try:
#统一转换为字符串处理 # 统一转换为字符串处理
expire_str = str(expire_time_str).strip() expire_str = str(expire_time_str).strip()
if expire_str.isdigit(): if expire_str.isdigit():
expire_time = datetime.fromtimestamp(int(expire_str)) expire_time = datetime.fromtimestamp(int(expire_str))
else: else:
expire_time = datetime.strptime( expire_time = datetime.strptime(expire_str, "%Y-%m-%dT%H:%M:%SZ")
expire_str,
"%Y-%m-%dT%H:%M:%SZ"
)
self.expire_time = expire_time.timestamp() - 60 self.expire_time = expire_time.timestamp() - 60
except Exception as e: except Exception as e:
raise ValueError(f"无效的过期时间格式: {expire_str}") from e raise ValueError(f"无效的过期时间格式: {expire_str}") from e
else: else:
self.expire_time = None self.expire_time = None
if not self.token: if not self.token:
raise ValueError("无法获取有效的访问Token") raise ValueError("无法获取有效的访问Token")
@@ -142,8 +162,12 @@ class TTSProvider(TTSProviderBase):
# f"过期时间 {datetime.fromtimestamp(self.expire_time)} | " # f"过期时间 {datetime.fromtimestamp(self.expire_time)} | "
# f"剩余 {remaining:.2f}秒") # f"剩余 {remaining:.2f}秒")
return time.time() > self.expire_time return time.time() > self.expire_time
def generate_filename(self, extension=".wav"): def generate_filename(self, extension=".wav"):
return os.path.join(self.output_file, f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}") return os.path.join(
self.output_file,
f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
)
async def text_to_speak(self, text, output_file): async def text_to_speak(self, text, output_file):
if self._is_token_expired(): if self._is_token_expired():
@@ -158,21 +182,27 @@ class TTSProvider(TTSProviderBase):
"voice": self.voice, "voice": self.voice,
"volume": self.volume, "volume": self.volume,
"speech_rate": self.speech_rate, "speech_rate": self.speech_rate,
"pitch_rate": self.pitch_rate "pitch_rate": self.pitch_rate,
} }
# print(self.api_url, json.dumps(request_json, ensure_ascii=False)) # print(self.api_url, json.dumps(request_json, ensure_ascii=False))
try: try:
resp = requests.post(self.api_url, json.dumps(request_json), headers=self.header) resp = requests.post(
self.api_url, json.dumps(request_json), headers=self.header
)
if resp.status_code == 401: # Token过期特殊处理 if resp.status_code == 401: # Token过期特殊处理
self._refresh_token() self._refresh_token()
resp = requests.post(self.api_url, json.dumps(request_json), headers=self.header) resp = requests.post(
self.api_url, json.dumps(request_json), headers=self.header
)
# 检查返回请求数据的mime类型是否是audio/***,是则保存到指定路径下;返回的是binary格式的 # 检查返回请求数据的mime类型是否是audio/***,是则保存到指定路径下;返回的是binary格式的
if resp.headers['Content-Type'].startswith('audio/'): if resp.headers["Content-Type"].startswith("audio/"):
with open(output_file, 'wb') as f: with open(output_file, "wb") as f:
f.write(resp.content) f.write(resp.content)
return output_file return output_file
else: else:
raise Exception(f"{__name__} status_code: {resp.status_code} response: {resp.content}") raise Exception(
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
)
except Exception as e: except Exception as e:
raise Exception(f"{__name__} error: {e}") raise Exception(f"{__name__} error: {e}")