diff --git a/README.md b/README.md index febc1f6e..1ec7b49d 100644 --- a/README.md +++ b/README.md @@ -235,7 +235,7 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/ --- ## 功能清单 ✨ ### 已实现 ✅ - +![请参考-全模块安装架构图](docs/images/deploy2.png) | 功能模块 | 描述 | |:---:|:---| | 核心架构 | 基于WebSocket和HTTP服务器,提供完整的控制台管理和认证系统 | @@ -271,7 +271,6 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/ --- ## 本项目支持的平台/组件列表 📋 - ### LLM 语言模型 | 使用方式 | 支持平台 | 免费平台 | diff --git a/README_en.md b/README_en.md index 006d1543..38b63f58 100644 --- a/README_en.md +++ b/README_en.md @@ -233,7 +233,7 @@ This project provides the following testing tools to help you verify the system --- ## Feature List ✨ ### Implemented ✅ - +![请参考-全模块安装架构图](docs/images/deploy2.png) | Feature Module | Description | |:---:|:---| | Core Architecture | Based on WebSocket and HTTP servers, provides complete console management and authentication system | @@ -269,7 +269,6 @@ Xiaozhi is an ecosystem. When using this product, you can also check out other e --- ## Supported Platforms/Components List 📋 - ### LLM Language Models | Usage Method | Supported Platforms | Free Platforms | diff --git a/docs/images/deploy2.png b/docs/images/deploy2.png index 61c36997..fa3f8de3 100644 Binary files a/docs/images/deploy2.png and b/docs/images/deploy2.png differ diff --git a/main/manager-api/src/main/java/xiaozhi/common/constant/Constant.java b/main/manager-api/src/main/java/xiaozhi/common/constant/Constant.java index f60a5fbe..b43099c8 100644 --- a/main/manager-api/src/main/java/xiaozhi/common/constant/Constant.java +++ b/main/manager-api/src/main/java/xiaozhi/common/constant/Constant.java @@ -237,7 +237,7 @@ public interface Constant { /** * 版本号 */ - public static final String VERSION = "0.7.2"; + public static final String VERSION = "0.7.3"; /** * 无效固件URL diff --git a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentServiceImpl.java index 36aa1192..96a4d014 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentServiceImpl.java @@ -41,6 +41,7 @@ import xiaozhi.modules.agent.service.AgentTemplateService; import xiaozhi.modules.agent.vo.AgentInfoVO; import xiaozhi.modules.device.service.DeviceService; import xiaozhi.modules.model.dto.ModelProviderDTO; +import xiaozhi.modules.model.entity.ModelConfigEntity; import xiaozhi.modules.model.service.ModelConfigService; import xiaozhi.modules.model.service.ModelProviderService; import xiaozhi.modules.security.user.SecurityUser; @@ -324,9 +325,32 @@ public class AgentServiceImpl extends BaseServiceImpl imp // 删除音频数据 agentChatHistoryService.deleteByAgentId(existingEntity.getId(), true, false); } + + boolean b = validateLLMIntentParams(dto.getLlmModelId(), dto.getIntentModelId()); + if (!b) { + throw new RenException("LLM大模型和Intent意图识别,选择参数不匹配"); + } this.updateById(existingEntity); } + /** + * 验证大语言模型和意图识别的参数是否符合匹配 + * + * @param llmModelId 大语言模型id + * @param intentModelId 意图识别id + * @return T 匹配 : F 不匹配 + */ + private boolean validateLLMIntentParams(String llmModelId, String intentModelId) { + ModelConfigEntity llmModelData = modelConfigService.selectById(llmModelId); + String type = llmModelData.getConfigJson().get("type").toString(); + // 如果查询大语言模型是openai或者ollama,意图识别选参数都可以 + if ("openai".equals(type) || "ollama".equals(type)) { + return true; + } + // 除了openai和ollama的类型,不可以选择id为Intent_function_call(函数调用)的意图识别 + return !"Intent_function_call".equals(intentModelId); + } + @Override @Transactional(rollbackFor = Exception.class) public String createAgent(AgentCreateDTO dto) { diff --git a/main/manager-api/src/main/java/xiaozhi/modules/model/controller/ModelController.java b/main/manager-api/src/main/java/xiaozhi/modules/model/controller/ModelController.java index 8baff29d..32a5c0cc 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/model/controller/ModelController.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/model/controller/ModelController.java @@ -21,11 +21,7 @@ import xiaozhi.common.utils.ConvertUtils; import xiaozhi.common.utils.Result; import xiaozhi.modules.agent.service.AgentTemplateService; import xiaozhi.modules.config.service.ConfigService; -import xiaozhi.modules.model.dto.ModelBasicInfoDTO; -import xiaozhi.modules.model.dto.ModelConfigBodyDTO; -import xiaozhi.modules.model.dto.ModelConfigDTO; -import xiaozhi.modules.model.dto.ModelProviderDTO; -import xiaozhi.modules.model.dto.VoiceDTO; +import xiaozhi.modules.model.dto.*; import xiaozhi.modules.model.entity.ModelConfigEntity; import xiaozhi.modules.model.service.ModelConfigService; import xiaozhi.modules.model.service.ModelProviderService; @@ -52,6 +48,14 @@ public class ModelController { return new Result>().ok(modelList); } + @GetMapping("/llm/names") + @Operation(summary = "获取LLM模型信息") + @RequiresPermissions("sys:role:normal") + public Result> getLlmModelCodeList(@RequestParam(required = false) String modelName) { + List llmModelCodeList = modelConfigService.getLlmModelCodeList(modelName); + return new Result>().ok(llmModelCodeList); + } + @GetMapping("/{modelType}/provideTypes") @Operation(summary = "获取模型供应器列表") @RequiresPermissions("sys:role:superAdmin") diff --git a/main/manager-api/src/main/java/xiaozhi/modules/model/dto/LlmModelBasicInfoDTO.java b/main/manager-api/src/main/java/xiaozhi/modules/model/dto/LlmModelBasicInfoDTO.java new file mode 100644 index 00000000..3995ceb8 --- /dev/null +++ b/main/manager-api/src/main/java/xiaozhi/modules/model/dto/LlmModelBasicInfoDTO.java @@ -0,0 +1,13 @@ +package xiaozhi.modules.model.dto; + +import lombok.Data; +import lombok.EqualsAndHashCode; + +/** + * LLM的模型的基础展示数据 + */ +@EqualsAndHashCode(callSuper = true) +@Data +public class LlmModelBasicInfoDTO extends ModelBasicInfoDTO{ + private String type; +} \ No newline at end of file diff --git a/main/manager-api/src/main/java/xiaozhi/modules/model/service/ModelConfigService.java b/main/manager-api/src/main/java/xiaozhi/modules/model/service/ModelConfigService.java index 9d15ac6d..634101e3 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/model/service/ModelConfigService.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/model/service/ModelConfigService.java @@ -4,6 +4,7 @@ import java.util.List; import xiaozhi.common.page.PageData; import xiaozhi.common.service.BaseService; +import xiaozhi.modules.model.dto.LlmModelBasicInfoDTO; import xiaozhi.modules.model.dto.ModelBasicInfoDTO; import xiaozhi.modules.model.dto.ModelConfigBodyDTO; import xiaozhi.modules.model.dto.ModelConfigDTO; @@ -13,6 +14,8 @@ public interface ModelConfigService extends BaseService { List getModelCodeList(String modelType, String modelName); + List getLlmModelCodeList(String modelName); + PageData getPageList(String modelType, String modelName, String page, String limit); ModelConfigDTO add(String modelType, String provideCode, ModelConfigBodyDTO modelConfigBodyDTO); diff --git a/main/manager-api/src/main/java/xiaozhi/modules/model/service/impl/ModelConfigServiceImpl.java b/main/manager-api/src/main/java/xiaozhi/modules/model/service/impl/ModelConfigServiceImpl.java index 1f006d93..c156b408 100644 --- a/main/manager-api/src/main/java/xiaozhi/modules/model/service/impl/ModelConfigServiceImpl.java +++ b/main/manager-api/src/main/java/xiaozhi/modules/model/service/impl/ModelConfigServiceImpl.java @@ -8,6 +8,7 @@ import java.util.stream.Collectors; import org.apache.commons.lang3.StringUtils; import org.springframework.stereotype.Service; +import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper; import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper; import com.baomidou.mybatisplus.core.metadata.IPage; @@ -23,6 +24,7 @@ import xiaozhi.common.utils.ConvertUtils; import xiaozhi.modules.agent.dao.AgentDao; import xiaozhi.modules.agent.entity.AgentEntity; import xiaozhi.modules.model.dao.ModelConfigDao; +import xiaozhi.modules.model.dto.LlmModelBasicInfoDTO; import xiaozhi.modules.model.dto.ModelBasicInfoDTO; import xiaozhi.modules.model.dto.ModelConfigBodyDTO; import xiaozhi.modules.model.dto.ModelConfigDTO; @@ -52,6 +54,25 @@ public class ModelConfigServiceImpl extends BaseServiceImpl getLlmModelCodeList(String modelName) { + List entities = modelConfigDao.selectList( + new QueryWrapper() + .eq("model_type", "llm") + .eq("is_enabled", 1) + .like(StringUtils.isNotBlank(modelName), "model_name", "%" + modelName + "%") + .select("id", "model_name", "config_json")); + // 处理获取到的内容 + return entities.stream().map(item -> { + LlmModelBasicInfoDTO dto = new LlmModelBasicInfoDTO(); + dto.setId(item.getId()); + dto.setModelName(item.getModelName()); + String type = item.getConfigJson().get("type").toString(); + dto.setType(type); + return dto; + }).toList(); + } + @Override public PageData getPageList(String modelType, String modelName, String page, String limit) { Map params = new HashMap(); @@ -94,6 +115,21 @@ public class ModelConfigServiceImpl extends BaseServiceImpl() + .eq(ModelConfigEntity::getId, llm)); + String selectModelType = (modelConfigEntity == null || modelConfigEntity.getModelType() == null) ? null + : modelConfigEntity.getModelType().toUpperCase(); + if (modelConfigEntity == null || !"LLM".equals(selectModelType)) { + throw new RenException("设置的LLM不存在"); + } + String type = modelConfigEntity.getConfigJson().get("type").toString(); + // 如果查询大语言模型是openai或者ollama,意图识别选参数都可以 + if (!"openai".equals(type) && !"ollama".equals(type)) { + throw new RenException("设置的LLM不是openai和ollama"); + } + } // 再更新供应器提供的模型 ModelConfigEntity modelConfigEntity = ConvertUtils.sourceToTarget(modelConfigBodyDTO, ModelConfigEntity.class); @@ -137,6 +173,8 @@ public class ModelConfigServiceImpl extends BaseServiceImpl { + RequestService.clearRequestTime(); + callback(res); + }) + .networkFail(() => { + RequestService.reAjaxFun(() => { + this.getLlmModelCodeList(modelName, callback); + }); + }).send(); + }, // 获取模型音色列表 getModelVoices(modelId, voiceName, callback) { const queryParams = new URLSearchParams({ diff --git a/main/manager-web/src/views/roleConfig.vue b/main/manager-web/src/views/roleConfig.vue index ccd96e1b..bbba01bd 100644 --- a/main/manager-web/src/views/roleConfig.vue +++ b/main/manager-web/src/views/roleConfig.vue @@ -89,7 +89,7 @@
-
@@ -130,7 +130,6 @@
- @@ -173,8 +172,9 @@ export default { { label: '视觉大模型(VLLM)', key: 'vllmModelId', type: 'VLLM' }, { label: '意图识别(Intent)', key: 'intentModelId', type: 'Intent' }, { label: '记忆(Memory)', key: 'memModelId', type: 'Memory' }, - { label: '语音合成(TTS)', key: 'ttsModelId', type: 'TTS' }, + { label: '语音合成(TTS)', key: 'ttsModelId', type: 'TTS' } ], + llmModeTypeMap: new Map(), modelOptions: {}, templates: [], loadingTemplate: false, @@ -356,6 +356,9 @@ export default { }); // 备份原始,以备取消时恢复 this.originalFunctions = JSON.parse(JSON.stringify(this.currentFunctions)); + + // 确保意图识别选项的可见性正确 + this.updateIntentOptionsVisibility(); }); } else { this.$message.error(data.msg || '获取配置失败'); @@ -364,16 +367,41 @@ export default { }, fetchModelOptions() { this.models.forEach(model => { - Api.model.getModelNames(model.type, '', ({ data }) => { - if (data.code === 0) { - this.$set(this.modelOptions, model.type, data.data.map(item => ({ - value: item.id, - label: item.modelName - }))); - } else { - this.$message.error(data.msg || '获取模型列表失败'); - } - }); + if (model.type != "LLM") { + Api.model.getModelNames(model.type, '', ({ data }) => { + if (data.code === 0) { + this.$set(this.modelOptions, model.type, data.data.map(item => ({ + value: item.id, + label: item.modelName, + isHidden: false + }))); + + // 如果是意图识别选项,需要根据当前LLM类型更新可见性 + if (model.type === 'Intent') { + this.updateIntentOptionsVisibility(); + } + } else { + this.$message.error(data.msg || '获取模型列表失败'); + } + }); + } else { + Api.model.getLlmModelCodeList('', ({ data }) => { + if (data.code === 0) { + let LLMdata = [] + data.data.forEach(item => { + LLMdata.push({ + value: item.id, + label: item.modelName, + isHidden: false + }) + this.llmModeTypeMap.set(item.id, item.type) + }) + this.$set(this.modelOptions, model.type, LLMdata); + } else { + this.$message.error(data.msg || '获取LLM模型列表失败'); + } + }); + } }); }, fetchVoiceOptions(modelId) { @@ -410,6 +438,10 @@ export default { if (type === 'Memory' && value !== 'Memory_nomem' && (this.form.chatHistoryConf === 0 || this.form.chatHistoryConf === null)) { this.form.chatHistoryConf = 2; } + if (type === 'LLM') { + // 当LLM类型改变时,更新意图识别选项的可见性 + this.updateIntentOptionsVisibility(); + } }, fetchAllFunctions() { return new Promise((resolve, reject) => { @@ -450,6 +482,42 @@ export default { } this.showFunctionDialog = false; }, + updateIntentOptionsVisibility() { + // 根据当前选择的LLM类型更新意图识别选项的可见性 + const currentLlmId = this.form.model.llmModelId; + if (!currentLlmId || !this.modelOptions['Intent']) return; + + const llmType = this.llmModeTypeMap.get(currentLlmId); + if (!llmType) return; + + this.modelOptions['Intent'].forEach(item => { + if (item.value === "Intent_function_call") { + // 如果llmType是openai或ollama,允许选择function_call + // 否则隐藏function_call选项 + if (llmType === "openai" || llmType === "ollama") { + item.isHidden = false; + } else { + item.isHidden = true; + } + } else { + // 其他意图识别选项始终可见 + item.isHidden = false; + } + }); + + // 如果当前选择的意图识别是function_call,但LLM类型不支持,则设置为可选的第一项 + if (this.form.model.intentModelId === "Intent_function_call" && + llmType !== "openai" && llmType !== "ollama") { + // 找到第一个可见的选项 + const firstVisibleOption = this.modelOptions['Intent'].find(item => !item.isHidden); + if (firstVisibleOption) { + this.form.model.intentModelId = firstVisibleOption.value; + } else { + // 如果没有可见选项,设置为Intent_nointent + this.form.model.intentModelId = 'Intent_nointent'; + } + } + }, updateChatHistoryConf() { if (this.form.model.memModelId === 'Memory_nomem') { this.form.chatHistoryConf = 0; diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml index 57f35603..8e6b2790 100644 --- a/main/xiaozhi-server/config.yaml +++ b/main/xiaozhi-server/config.yaml @@ -710,6 +710,26 @@ TTS: # voice_id: female-shaonv # weight: 1 # language_boost: auto + +# MinimaxTTSHTTPStream和MinimaxTTSWebSocketStream还在测试,测试完再开放 +# +# MinimaxTTSHTTPStream: +# # Minimax流式语音合成服务 +# type: minimax_httpstream +# output_dir: tmp/ +# group_id: 你的minimax平台groupID +# api_key: 你的minimax平台接口密钥 +# model: "speech-01-turbo" +# voice_id: "female-shaonv" +# +# MinimaxTTSWebSocketStream: +# type: minimax_webSocket +# output_dir: tmp/ +# group_id: 你的minimax平台groupID +# api_key: 你的minimax平台接口密钥 +# model: "speech-01-turbo" +# voice_id: "female-shaonv" + AliyunTTS: # 阿里云智能语音交互服务,需要先在阿里云平台开通服务,然后获取验证信息 # 平台地址:https://nls-portal.console.aliyun.com/ diff --git a/main/xiaozhi-server/config/logger.py b/main/xiaozhi-server/config/logger.py index 7640f17d..95577e7e 100644 --- a/main/xiaozhi-server/config/logger.py +++ b/main/xiaozhi-server/config/logger.py @@ -5,7 +5,7 @@ from config.config_loader import load_config from config.settings import check_config_file from datetime import datetime -SERVER_VERSION = "0.7.2" +SERVER_VERSION = "0.7.3" _logger_initialized = False diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index fba87232..7221281f 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -132,6 +132,8 @@ class ConnectionHandler: # tts相关变量 self.sentence_id = None + # 处理TTS响应没有文本返回 + self.tts_MessageText = "" # iot相关变量 self.iot_descriptors = {} @@ -182,8 +184,13 @@ class ConnectionHandler: await ws.send("端口正常,如需测试连接,请使用test_page.html") await self.close(ws) return - # 获取客户端ip地址 - self.client_ip = ws.remote_address[0] + real_ip = self.headers.get("x-real-ip") or self.headers.get( + "x-forwarded-for" + ) + if real_ip: + self.client_ip = real_ip.split(",")[0].strip() + else: + self.client_ip = ws.remote_address[0] self.logger.bind(tag=TAG).info( f"{self.client_ip} conn - Headers: {self.headers}" ) @@ -335,7 +342,7 @@ class ConnectionHandler: self.config.get("selected_module", {}) ) self.logger = create_connection_logger(self.selected_module_str) - + """初始化组件""" if self.config.get("prompt") is not None: user_prompt = self.config["prompt"] @@ -351,10 +358,10 @@ class ConnectionHandler: self.vad = self._vad if self.asr is None: self.asr = self._initialize_asr() - + # 初始化声纹识别 self._initialize_voiceprint() - + # 打开语音识别通道 asyncio.run_coroutine_threadsafe( self.asr.open_audio_channels(self), self.loop @@ -746,7 +753,7 @@ class ConnectionHandler: content = response # 在llm回复中获取情绪表情,一轮对话只在开头获取一次 - if emotion_flag: + if emotion_flag and content is not None and content.strip(): asyncio.run_coroutine_threadsafe( textUtils.get_emotion(self, content), self.loop, @@ -790,9 +797,9 @@ class ConnectionHandler: if not bHasError: # 如需要大模型先处理一轮,添加相关处理后的日志情况 if len(response_message) > 0: - self.dialogue.put( - Message(role="assistant", content="".join(response_message)) - ) + text_buff = "".join(response_message) + self.tts_MessageText = text_buff + self.dialogue.put(Message(role="assistant", content=text_buff)) response_message.clear() self.logger.bind(tag=TAG).debug( f"function_name={function_name}, function_id={function_id}, function_arguments={function_arguments}" @@ -814,9 +821,9 @@ class ConnectionHandler: # 存储对话内容 if len(response_message) > 0: - self.dialogue.put( - Message(role="assistant", content="".join(response_message)) - ) + text_buff = "".join(response_message) + self.tts_MessageText = text_buff + self.dialogue.put(Message(role="assistant", content=text_buff)) if depth == 0: self.tts.tts_text_queue.put( TTSMessageDTO( @@ -893,9 +900,7 @@ class ConnectionHandler: if self.executor is None: continue # 提交任务到线程池 - self.executor.submit( - self._process_report, *item - ) + self.executor.submit(self._process_report, *item) except Exception as e: self.logger.bind(tag=TAG).error(f"聊天记录上报线程异常: {e}") except queue.Empty: diff --git a/main/xiaozhi-server/core/handle/sendAudioHandle.py b/main/xiaozhi-server/core/handle/sendAudioHandle.py index 2430446c..aa59a4d1 100644 --- a/main/xiaozhi-server/core/handle/sendAudioHandle.py +++ b/main/xiaozhi-server/core/handle/sendAudioHandle.py @@ -12,7 +12,7 @@ async def sendAudioMessage(conn, sentenceType, audios, text): conn.logger.bind(tag=TAG).info(f"发送音频消息: {sentenceType}, {text}") pre_buffer = False - if conn.tts.tts_audio_first_sentence and text is not None: + if conn.tts.tts_audio_first_sentence: conn.logger.bind(tag=TAG).info(f"发送第一段语音: {text}") conn.tts.tts_audio_first_sentence = False pre_buffer = True @@ -73,7 +73,7 @@ async def send_tts_message(conn, state, text=None): """发送 TTS 状态消息""" message = {"type": "tts", "state": state, "session_id": conn.session_id} if text is not None: - message["text"] = text + message["text"] = textUtils.check_emoji(text) # TTS播放结束 if state == "stop": diff --git a/main/xiaozhi-server/core/providers/tts/aliyun_stream.py b/main/xiaozhi-server/core/providers/tts/aliyun_stream.py index 2a2a6891..0131fcfd 100644 --- a/main/xiaozhi-server/core/providers/tts/aliyun_stream.py +++ b/main/xiaozhi-server/core/providers/tts/aliyun_stream.py @@ -127,8 +127,6 @@ class TTSProvider(TTSProviderBase): # 专属tts设置 self.message_id = "" - self.tts_text = "" - self.text_buffer = [] # 创建Opus编码器 self.opus_encoder = opus_encoder_utils.OpusEncoderUtils( @@ -229,7 +227,6 @@ class TTSProvider(TTSProviderBase): # aliyunStream独有的参数生成 self.message_id = str(uuid.uuid4().hex) - self.text_buffer = [] logger.bind(tag=TAG).info("开始启动TTS会话...") future = asyncio.run_coroutine_threadsafe( @@ -250,7 +247,6 @@ class TTSProvider(TTSProviderBase): logger.bind(tag=TAG).debug( f"开始发送TTS文本: {message.content_detail}" ) - self.text_buffer.append(message.content_detail) future = asyncio.run_coroutine_threadsafe( self.text_to_speak(message.content_detail, None), loop=self.conn.loop, @@ -275,9 +271,6 @@ class TTSProvider(TTSProviderBase): if message.sentence_type == SentenceType.LAST: try: logger.bind(tag=TAG).info("开始结束TTS会话...") - self.tts_text = textUtils.get_string_no_punctuation_or_emoji( - "".join(self.text_buffer).replace("\n", "") - ) future = asyncio.run_coroutine_threadsafe( self.finish_session(self.conn.sentence_id), loop=self.conn.loop, @@ -444,34 +437,35 @@ class TTSProvider(TTSProviderBase): event_name = header.get("name") if event_name == "SynthesisStarted": logger.bind(tag=TAG).debug("TTS合成已启动") - elif event_name == "SentenceBegin": - logger.bind(tag=TAG).debug( - f"句子语音生成开始: {self.tts_text}" - ) - opus_datas_cache = [] self.tts_audio_queue.put( - (SentenceType.FIRST, [], self.tts_text) + (SentenceType.FIRST, [], None) ) + elif event_name == "SentenceBegin": + opus_datas_cache = [] elif event_name == "SentenceEnd": - logger.bind(tag=TAG).info( - f"句子语音生成成功: {self.tts_text}" - ) if ( not is_first_sentence or first_sentence_segment_count > 10 ): # 发送缓存的数据 - self.tts_audio_queue.put( - (SentenceType.MIDDLE, opus_datas_cache, None) - ) + if self.conn.tts_MessageText: + logger.bind(tag=TAG).info( + f"句子语音生成成功: {self.conn.tts_MessageText}" + ) + self.tts_audio_queue.put( + (SentenceType.MIDDLE, opus_datas_cache, self.conn.tts_MessageText) + ) + self.conn.tts_MessageText = None + else: + self.tts_audio_queue.put( + (SentenceType.MIDDLE, opus_datas_cache, None) + ) # 第一句话结束后,将标志设置为False is_first_sentence = False elif event_name == "SynthesisCompleted": logger.bind(tag=TAG).debug(f"会话结束~~") self._process_before_stop_play_files() session_finished = True - self.reuse_judgment = time.time() - self.tts_text = "" break except json.JSONDecodeError: logger.bind(tag=TAG).warning("收到无效的JSON消息") diff --git a/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py b/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py index 82cba035..ea7f18fd 100644 --- a/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py +++ b/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py @@ -232,7 +232,6 @@ class TTSProvider(TTSProviderBase): loop=self.conn.loop, ) future.result() - self.tts_audio_first_sentence = True self.before_stop_play_files.clear() logger.bind(tag=TAG).info("TTS会话启动成功") except Exception as e: diff --git a/main/xiaozhi-server/core/providers/tts/linkerai.py b/main/xiaozhi-server/core/providers/tts/linkerai.py index 1046b972..d1be586c 100644 --- a/main/xiaozhi-server/core/providers/tts/linkerai.py +++ b/main/xiaozhi-server/core/providers/tts/linkerai.py @@ -52,7 +52,6 @@ class TTSProvider(TTSProviderBase): self.processed_chars = 0 self.tts_text_buff = [] self.segment_count = 0 - self.tts_audio_first_sentence = True self.before_stop_play_files.clear() elif ContentType.TEXT == message.content_type: self.tts_text_buff.append(message.content_detail) diff --git a/main/xiaozhi-server/core/providers/tts/minimax_httpstream.py b/main/xiaozhi-server/core/providers/tts/minimax_httpstream.py new file mode 100644 index 00000000..605b3efd --- /dev/null +++ b/main/xiaozhi-server/core/providers/tts/minimax_httpstream.py @@ -0,0 +1,230 @@ +import os +import uuid +import json +import requests +from datetime import datetime +from typing import Iterator, Optional, Union +from core.providers.tts.base import TTSProviderBase +from core.utils.util import parse_string_to_list + + +class TTSProvider(TTSProviderBase): + def __init__(self, config, delete_audio_file): + super().__init__(config, delete_audio_file) + self.group_id = config.get("group_id") + self.api_key = config.get("api_key") + self.model = config.get("model") + if config.get("private_voice"): + self.voice = config.get("private_voice") + else: + self.voice = config.get("voice_id") + + default_voice_setting = { + "voice_id": "female-shaonv", + "speed": 1, + "vol": 1, + "pitch": 0, + "emotion": "happy", + } + default_pronunciation_dict = {"tone": ["处理/(chu3)(li3)", "危险/dangerous"]} + defult_audio_setting = { + "sample_rate": 32000, + "bitrate": 128000, + "format": "mp3", + "channel": 1, + } + self.voice_setting = { + **default_voice_setting, + **config.get("voice_setting", {}), + } + self.pronunciation_dict = { + **default_pronunciation_dict, + **config.get("pronunciation_dict", {}), + } + self.audio_setting = {**defult_audio_setting, **config.get("audio_setting", {})} + self.timber_weights = parse_string_to_list(config.get("timber_weights")) + + if self.voice: + self.voice_setting["voice_id"] = self.voice + + self.host = "api.minimax.chat" + self.api_url = f"https://{self.host}/v1/t2a_v2?GroupId={self.group_id}" + self.header = { + "Content-Type": "application/json", + "Authorization": f"Bearer {self.api_key}", + } + self.audio_file_type = defult_audio_setting.get("format", "mp3") + + def generate_filename(self, extension=".mp3"): + 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): + """非流式语音合成(保留原有实现)""" + request_json = { + "model": self.model, + "text": text, + "stream": False, + "voice_setting": self.voice_setting, + "pronunciation_dict": self.pronunciation_dict, + "audio_setting": self.audio_setting, + } + + if type(self.timber_weights) is list and len(self.timber_weights) > 0: + request_json["timber_weights"] = self.timber_weights + request_json["voice_setting"]["voice_id"] = "" + + try: + resp = requests.post( + self.api_url, json.dumps(request_json), headers=self.header + ) + if resp.json()["base_resp"]["status_code"] == 0: + data = resp.json()["data"]["audio"] + audio_bytes = bytes.fromhex(data) + if output_file: + with open(output_file, "wb") as file_to_save: + file_to_save.write(audio_bytes) + else: + return audio_bytes + else: + raise Exception( + f"{__name__} status_code: {resp.status_code} response: {resp.content}" + ) + except Exception as e: + raise Exception(f"{__name__} error: {e}") + + def text_to_speak_stream( + self, + text: str, + chunk_callback: Optional[callable] = None + ) -> Iterator[bytes]: + """ + 流式语音合成方法 + :param text: 要合成的文本 + :param chunk_callback: 可选的回调函数,用于处理每个音频块 + :return: 生成器,每次产生一个音频数据块(bytes) + """ + request_json = { + "model": self.model, + "text": text, + "stream": True, + "voice_setting": self.voice_setting, + "pronunciation_dict": self.pronunciation_dict, + "audio_setting": self.audio_setting, + } + + if isinstance(self.timber_weights, list) and len(self.timber_weights) > 0: + request_json["timber_weights"] = self.timber_weights + request_json["voice_setting"]["voice_id"] = "" + + try: + with requests.post( + self.api_url, + data=json.dumps(request_json), + headers=self.header, + stream=True + ) as response: + + # 检查HTTP状态码 + if response.status_code != 200: + raise Exception( + f"HTTP error: {response.status_code}, response: {response.text}" + ) + + # 处理流式响应 + for line in response.iter_lines(): + if line: # 过滤空行 + # 检查是否为数据行 (SSE格式) + if line.startswith(b'data:'): + try: + data = json.loads(line[5:].strip()) # 去掉"data:"前缀 + + # 检查API状态码 + if data.get("base_resp", {}).get("status_code", -1) != 0: + raise Exception( + f"API error: {data.get('base_resp', {}).get('status_msg')}" + ) + + # 跳过非音频数据块 + if "extra_info" in data: + continue + + # 提取音频数据 + audio_hex = data.get("data", {}).get("audio") + if audio_hex: + audio_chunk = bytes.fromhex(audio_hex) + if chunk_callback: + chunk_callback(audio_chunk) + yield audio_chunk + + except json.JSONDecodeError: + # 忽略JSON解析错误(可能是心跳包等) + continue + except Exception as e: + raise e + + except Exception as e: + raise Exception(f"{__name__} stream error: {e}") + + def save_stream_to_file( + self, + text: str, + output_file: Optional[str] = None, + progress_callback: Optional[callable] = None + ) -> str: + """ + 流式合成并保存到文件 + :param text: 要合成的文本 + :param output_file: 输出文件路径,如果为None则自动生成 + :param progress_callback: 可选的回调函数,接收已写入的字节数 + :return: 保存的文件路径 + """ + if not output_file: + output_file = self.generate_filename(extension=f".{self.audio_file_type}") + + os.makedirs(os.path.dirname(output_file), exist_ok=True) + + total_bytes = 0 + try: + with open(output_file, "wb") as audio_file: + for audio_chunk in self.text_to_speak_stream(text): + audio_file.write(audio_chunk) + audio_file.flush() + total_bytes += len(audio_chunk) + if progress_callback: + progress_callback(total_bytes) + return output_file + except Exception as e: + # 清理可能创建的不完整文件 + if os.path.exists(output_file): + os.remove(output_file) + raise e + + def stream_to_audio_player(self, text: str, player_command: list = None): + """ + 流式合成并直接播放音频 + :param text: 要合成的文本 + :param player_command: 音频播放器命令,默认使用mpv + """ + if player_command is None: + player_command = ["mpv", "--no-cache", "--no-terminal", "--", "fd://0"] + + try: + import subprocess + player_process = subprocess.Popen( + player_command, + stdin=subprocess.PIPE, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + ) + + for audio_chunk in self.text_to_speak_stream(text): + player_process.stdin.write(audio_chunk) + player_process.stdin.flush() + + player_process.stdin.close() + player_process.wait() + except Exception as e: + raise Exception(f"Audio player error: {e}") \ No newline at end of file diff --git a/main/xiaozhi-server/core/providers/tts/minimax_webSocket.py b/main/xiaozhi-server/core/providers/tts/minimax_webSocket.py new file mode 100644 index 00000000..bda6b47f --- /dev/null +++ b/main/xiaozhi-server/core/providers/tts/minimax_webSocket.py @@ -0,0 +1,180 @@ +import os +import uuid +import json +import asyncio +import websockets +import ssl +from datetime import datetime +from core.providers.tts.base import TTSProviderBase +from core.utils.util import parse_string_to_list + + +class TTSProvider(TTSProviderBase): + def __init__(self, config, delete_audio_file): + super().__init__(config, delete_audio_file) + self.group_id = config.get("group_id") + self.api_key = config.get("api_key") + self.model = config.get("model") + + # 初始化语音设置 + default_voice_setting = { + "voice_id": "female-shaonv", + "speed": 1, + "vol": 1, + "pitch": 0, + "emotion": "happy", + } + default_pronunciation_dict = {"tone": ["处理/(chu3)(li3)", "危险/dangerous"]} + default_audio_setting = { + "sample_rate": 32000, + "bitrate": 128000, + "format": "mp3", + "channel": 1, + } + + # 合并配置 + self.voice_setting = { + **default_voice_setting, + **config.get("voice_setting", {}), + } + self.pronunciation_dict = { + **default_pronunciation_dict, + **config.get("pronunciation_dict", {}), + } + self.audio_setting = { + **default_audio_setting, + **config.get("audio_setting", {}) + } + self.timber_weights = parse_string_to_list(config.get("timber_weights")) + + # 设置语音ID + if config.get("private_voice"): + self.voice_setting["voice_id"] = config.get("private_voice") + elif config.get("voice_id"): + self.voice_setting["voice_id"] = config.get("voice_id") + + # WebSocket配置 + self.ws_url = "wss://api.minimaxi.com/ws/v1/t2a_v2" + self.headers = { + "Authorization": f"Bearer {self.api_key}", + "GroupId": self.group_id + } + self.audio_file_type = self.audio_setting.get("format", "mp3") + + def generate_filename(self, extension=".mp3"): + """生成唯一的音频文件名""" + return os.path.join( + self.output_file, + f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}", + ) + + async def _establish_connection(self): + """建立WebSocket连接""" + ssl_context = ssl.create_default_context() + ssl_context.check_hostname = False + ssl_context.verify_mode = ssl.CERT_NONE + + try: + ws = await websockets.connect( + self.ws_url, + additional_headers=self.headers, + ssl=ssl_context + ) + connected = json.loads(await ws.recv()) + if connected.get("event") == "connected_success": + print("连接成功") + return ws + return None + except Exception as e: + print(f"连接失败: {e}") + return None + + async def _start_task(self, websocket): + """发送任务开始请求""" + start_msg = { + "event": "task_start", + "model": self.model, + "voice_setting": self.voice_setting, + "pronunciation_dict": self.pronunciation_dict, + "audio_setting": self.audio_setting + } + + if self.timber_weights and len(self.timber_weights) > 0: + start_msg["timber_weights"] = self.timber_weights + start_msg["voice_setting"]["voice_id"] = "" + + await websocket.send(json.dumps(start_msg)) + response = json.loads(await websocket.recv()) + return response.get("event") == "task_started" + + async def _continue_task(self, websocket, text): + """发送继续请求并收集音频数据""" + await websocket.send(json.dumps({ + "event": "task_continue", + "text": text + })) + + audio_chunks = [] + while True: + response = json.loads(await websocket.recv()) + if "data" in response and "audio" in response["data"]: + audio_chunks.append(response["data"]["audio"]) + if response.get("is_final"): + break + return "".join(audio_chunks) + + async def _close_connection(self, websocket): + """关闭连接""" + if websocket: + await websocket.send(json.dumps({"event": "task_finish"})) + await websocket.close() + print("连接已关闭") + + async def text_to_speak(self, text, output_file=None): + """主方法:文本转语音""" + ws = await self._establish_connection() + if not ws: + raise Exception("无法建立WebSocket连接") + + try: + if not await self._start_task(ws): + raise Exception("任务启动失败") + + hex_audio = await self._continue_task(ws, text) + audio_bytes = bytes.fromhex(hex_audio) + + # 保存到文件或返回二进制数据 + if output_file: + with open(output_file, "wb") as f: + f.write(audio_bytes) + print(f"音频已保存为{output_file}") + return output_file + else: + # 返回音频二进制数据(不播放) + return audio_bytes + + finally: + await self._close_connection(ws) + + +async def main(): + """测试用主函数""" + # 示例配置 + config = { + "group_id": "YOUR_GROUP_ID", # 替换为实际的group_id + "api_key": "YOUR_API_KEY", # 替换为实际的api_key + "model": "your-model", # 替换为实际的模型名称 + "voice_id": "male-qn-qingse", + "voice_setting": { + "speed": 1.2, + "emotion": "happy" + } + } + + tts = TTSProvider(config, delete_audio_file=True) + output_file = tts.generate_filename() + await tts.text_to_speak("这是一个测试文本,用于验证流式语音合成功能", output_file) + + +if __name__ == "__main__": + asyncio.run(main()) \ No newline at end of file diff --git a/main/xiaozhi-server/core/utils/textUtils.py b/main/xiaozhi-server/core/utils/textUtils.py index dd876cfa..47677b99 100644 --- a/main/xiaozhi-server/core/utils/textUtils.py +++ b/main/xiaozhi-server/core/utils/textUtils.py @@ -24,6 +24,15 @@ EMOJI_MAP = { "😘": "kissy", "😏": "confident", } +EMOJI_RANGES = [ + (0x1F600, 0x1F64F), + (0x1F300, 0x1F5FF), + (0x1F680, 0x1F6FF), + (0x1F900, 0x1F9FF), + (0x1FA70, 0x1FAFF), + (0x2600, 0x26FF), + (0x2700, 0x27BF), +] def get_string_no_punctuation_or_emoji(s): @@ -65,18 +74,7 @@ def is_punctuation_or_emoji(char): } if char.isspace() or char in punctuation_set: return True - # 检查表情符号(保留原有逻辑) - code_point = ord(char) - emoji_ranges = [ - (0x1F600, 0x1F64F), - (0x1F300, 0x1F5FF), - (0x1F680, 0x1F6FF), - (0x1F900, 0x1F9FF), - (0x1FA70, 0x1FAFF), - (0x2600, 0x26FF), - (0x2700, 0x27BF), - ] - return any(start <= code_point <= end for start, end in emoji_ranges) + return is_emoji(char) async def get_emotion(conn, text): @@ -102,3 +100,14 @@ async def get_emotion(conn, text): except Exception as e: conn.logger.bind(tag=TAG).warning(f"发送情绪表情失败,错误:{e}") return + + +def is_emoji(char): + """检查字符是否为emoji表情""" + code_point = ord(char) + return any(start <= code_point <= end for start, end in EMOJI_RANGES) + + +def check_emoji(text): + """去除文本中的所有emoji表情""" + return ''.join(char for char in text if not is_emoji(char) and char != "\n")