mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-26 17:13:54 +08:00
Merge pull request #1351 from xinnan-tech/py_memory_llm
update: 记忆模块使用独立LLM openai增加超参
This commit is contained in:
+2
@@ -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;
|
||||||
}
|
}
|
||||||
|
|||||||
+7
-5
@@ -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();
|
||||||
|
|
||||||
// 保存数据
|
// 保存数据
|
||||||
|
|||||||
+26
-4
@@ -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);
|
||||||
|
|||||||
@@ -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';
|
||||||
@@ -163,3 +163,10 @@ databaseChangeLog:
|
|||||||
- 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
|
||||||
@@ -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
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -16,38 +16,44 @@ class LLMProvider(LLMProviderBase):
|
|||||||
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,9 +65,13 @@ 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,
|
||||||
@@ -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,7 +85,6 @@ 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)
|
||||||
|
|
||||||
@@ -80,16 +94,27 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
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,28 +124,23 @@ 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
|
||||||
@@ -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}")
|
||||||
|
|||||||
Reference in New Issue
Block a user