Compare commits

..
Author SHA1 Message Date
CGDandGitHub de45f73efd Merge pull request #3303 from chentyke/fix/frontend-contract-consistency
fix(manager): align frontend and backend field contracts
2026-07-28 09:37:39 +08:00
CGDandGitHub 1dc46e30d7 Merge pull request #3302 from xinnan-tech/refactor/manager-api-tool-dependencies
refactor(manager-api): 升级常用工具依赖并收敛冗余实现
2026-07-28 09:28:00 +08:00
Sakura-RanChenandGitHub 5e0d853256 Merge pull request #3309 from xinnan-tech/py-fix-reportHandle
refactor(reportHandle): 重构上报处理逻辑,适配多格式音频转换
2026-07-27 15:24:15 +08:00
CGDandGitHub f5ed1aaec8 Merge pull request #3305 from xinnan-tech/update_version
Bump to 0.9.6
2026-07-24 11:32:33 +08:00
Tyke Chen 18e5c47a06 fix(manager): align frontend and backend field contracts 2026-07-24 10:21:05 +08:00
hrz 9da58e254c Bump to 0.9.6 2026-07-24 09:53:03 +08:00
Tyke Chen 4a2c90ec87 refactor(manager-api): 升级常用工具依赖并收敛冗余实现 2026-07-24 09:06:15 +08:00
CGDandGitHub 27e57631a7 Merge pull request #3300 from xinnan-tech/test/manager-api-runtime-warnings
test(manager-api): 清理测试运行告警
2026-07-23 17:46:04 +08:00
CGDandGitHub b1e39da9b8 Merge pull request #3301 from chentyke/fix/issue-3299-auto-update-state
fix: 修复设备自动升级开关状态不生效
2026-07-22 17:38:13 +08:00
Tyke Chen e87fc766d2 fix: restore device auto-update switch behavior 2026-07-22 17:10:21 +08:00
Tyke Chen 942d55118b test(manager-api): clean test runtime warnings 2026-07-22 16:41:44 +08:00
CGDandGitHub f47f3b050b Merge pull request #3297 from xinnan-tech/fix/manager-api-compiler-warnings
fix(manager-api): 清理编译器告警
2026-07-22 16:22:46 +08:00
wengzhandGitHub e3453b1a87 Merge pull request #3298 from xinnan-tech/py-fix-intent
Py fix intent
2026-07-21 15:16:19 +08:00
Tyke Chen d760465408 fix(manager-api): make JSON collections type-safe 2026-07-21 14:43:25 +08:00
Tyke Chen bdb2538239 fix(manager-api): clean deprecated and generic warnings 2026-07-21 14:43:25 +08:00
CGDandGitHub ea6c144e43 Merge pull request #3294 from xinnan-tech/fix/manager-api-startup-warnings
fix(manager-api): 消除启动阶段告警
2026-07-21 14:37:22 +08:00
CGDandGitHub a79aa455d6 Merge pull request #3292 from xinnan-tech/fix/manager-api-test-i18n
test: 修复 manager-api 国际化测试配置
2026-07-21 14:29:30 +08:00
wengzh 6e92d169ec refactor(get_news_from_newsnow): 改用流式处理解析网页内容
将原有的直接转换响应内容改为流式读取字节流并传入参数,适配MarkItDown的convert_stream接口,优化大内容处理时的内存占用
2026-07-20 15:51:23 +08:00
Sakura-RanChen ed89c05245 fix: intent阻塞主线程 2026-07-20 14:48:26 +08:00
Tyke Chen b500d1c6bd fix(manager-api): avoid duplicate address book insert mapping 2026-07-20 14:46:14 +08:00
Tyke Chen 0598b09629 fix(manager-api): remove invalid logback startup directives 2026-07-20 14:46:14 +08:00
Tyke Chen e4907121f4 fix(manager-api): eliminate premature bean initialization warnings 2026-07-20 14:46:14 +08:00
Tyke Chen 2a618d2f8f test: 修复 manager-api 国际化测试配置 2026-07-20 10:34:33 +08:00
81 changed files with 1401 additions and 418 deletions
+16 -4
View File
@@ -21,18 +21,19 @@
<junit.version>5.10.1</junit.version>
<druid.version>1.2.20</druid.version>
<mybatisplus.version>3.5.17</mybatisplus.version>
<hutool.version>5.8.24</hutool.version>
<jsoup.version>1.19.1</jsoup.version>
<hutool.version>5.8.46</hutool.version>
<jsoup.version>1.22.2</jsoup.version>
<knife4j.version>4.6.0</knife4j.version>
<springdoc.version>2.8.8</springdoc.version>
<commons-lang3.version>3.18.0</commons-lang3.version>
<commons-lang3.version>3.20.0</commons-lang3.version>
<shiro.version>2.0.2</shiro.version>
<captcha.version>1.6.2</captcha.version>
<guava.version>33.0.0-jre</guava.version>
<guava.version>33.6.0-jre</guava.version>
<liquibase-core.version>4.20.0</liquibase-core.version>
<aliyun-sms-version>4.1.0</aliyun-sms-version>
<okio-version>3.4.0</okio-version>
<skipTests>true</skipTests>
<argLine></argLine>
</properties>
<dependencies>
@@ -272,11 +273,22 @@
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-maven-plugin</artifactId>
</plugin>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-compiler-plugin</artifactId>
<configuration>
<proc>full</proc>
<compilerArgs>
<arg>-Xlint:deprecation,unchecked</arg>
</compilerArgs>
</configuration>
</plugin>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-surefire-plugin</artifactId>
<configuration>
<skipTests>${skipTests}</skipTests>
<argLine>@{argLine} -Xshare:off -javaagent:"${settings.localRepository}/org/mockito/mockito-core/${mockito.version}/mockito-core-${mockito.version}.jar"</argLine>
</configuration>
</plugin>
</plugins>
@@ -324,7 +324,7 @@ public interface Constant {
/**
* 版本号
*/
public static final String VERSION = "0.9.5";
public static final String VERSION = "0.9.6";
/**
* 无效固件URL
@@ -67,7 +67,7 @@ public class DataFilterInterceptor implements InnerInterceptor {
private String getSelect(String buildSql, DataScope scope) {
try {
Select select = (Select) CCJSqlParserUtil.parse(buildSql);
PlainSelect plainSelect = (PlainSelect) select.getSelectBody();
PlainSelect plainSelect = select.getPlainSelect();
Expression expression = plainSelect.getWhere();
if (expression == null) {
@@ -81,8 +81,8 @@ public abstract class BaseServiceImpl<M extends BaseMapper<T>, T> implements Bas
// 处理排序字段
if (orderField instanceof String) {
orderFields.add((String) orderField);
} else if (orderField instanceof List) {
orderFields.addAll((List<String>) orderField);
} else if (orderField instanceof List<?> fields) {
fields.forEach(field -> orderFields.add(String.class.cast(field)));
}
// 有排序字段则排序
@@ -142,11 +142,12 @@ public abstract class BaseServiceImpl<M extends BaseMapper<T>, T> implements Bas
return SqlHelper.retBool(result);
}
protected Class<M> currentMapperClass() {
return (Class<M>) ReflectionKit.getSuperClassGenericType(this.getClass(), BaseServiceImpl.class, 0);
protected Class<?> currentMapperClass() {
return ReflectionKit.getSuperClassGenericType(this.getClass(), BaseServiceImpl.class, 0);
}
@Override
@SuppressWarnings("unchecked")
public Class<T> currentModelClass() {
return (Class<T>) ReflectionKit.getSuperClassGenericType(this.getClass(), BaseServiceImpl.class, 1);
}
@@ -226,6 +227,6 @@ public abstract class BaseServiceImpl<M extends BaseMapper<T>, T> implements Bas
@Override
public boolean deleteBatchIds(Collection<? extends Serializable> idList) {
return SqlHelper.retBool(baseDao.deleteBatchIds(idList));
return SqlHelper.retBool(baseDao.deleteByIds(idList));
}
}
@@ -24,6 +24,7 @@ import xiaozhi.common.utils.ConvertUtils;
public abstract class CrudServiceImpl<M extends BaseMapper<T>, T, D> extends BaseServiceImpl<M, T>
implements CrudService<T, D> {
@SuppressWarnings("unchecked")
protected Class<D> currentDtoClass() {
return (Class<D>) ReflectionKit.getSuperClassGenericType(getClass(), CrudServiceImpl.class, 2);
}
@@ -70,6 +71,6 @@ public abstract class CrudServiceImpl<M extends BaseMapper<T>, T, D> extends Bas
@Override
public void delete(Serializable[] ids) {
baseDao.deleteBatchIds(Arrays.asList(ids));
baseDao.deleteByIds(Arrays.asList(ids));
}
}
@@ -1,52 +0,0 @@
package xiaozhi.common.utils;
import lombok.extern.slf4j.Slf4j;
import java.security.MessageDigest;
import java.security.NoSuchAlgorithmException;
/**
* 哈希加密算法的工具类
* @author zjy
*/
@Slf4j
public class HashEncryptionUtil {
/**
* 使用md5进行加密
* @param context 被加密的内容
* @return 哈希值
*/
public static String Md5hexDigest(String context){
return hexDigest(context,"MD5");
}
/**
* 指定哈希算法进行加密
* @param context 被加密的内容
* @param algorithm 哈希算法
* @return 哈希值
*/
public static String hexDigest(String context,String algorithm ){
// 获取MD5算法实例
MessageDigest md = null;
try {
md = MessageDigest.getInstance(algorithm);
} catch (NoSuchAlgorithmException e) {
log.error("加密失败的算法:{}",algorithm);
throw new RuntimeException("加密失败,"+ algorithm +"哈希算法系统不支持");
}
// 计算智能体id的MD5值
byte[] messageDigest = md.digest(context.getBytes());
// 将字节数组转换为十六进制字符串
StringBuilder hexString = new StringBuilder();
for (byte b : messageDigest) {
String hex = Integer.toHexString(0xFF & b);
if (hex.length() == 1) {
hexString.append('0');
}
hexString.append(hex);
}
return hexString.toString();
}
}
@@ -1,7 +1,9 @@
package xiaozhi.common.utils;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.ObjectMapper;
@@ -16,6 +18,10 @@ import cn.hutool.core.util.StrUtil;
*/
public class JsonUtils {
private static final ObjectMapper objectMapper = new ObjectMapper();
private static final TypeReference<Map<String, Object>> STRING_OBJECT_MAP = new TypeReference<>() {
};
private static final TypeReference<List<Map<String, Object>>> STRING_OBJECT_MAP_LIST = new TypeReference<>() {
};
public static String toJsonString(Object object) {
try {
@@ -67,4 +73,59 @@ public class JsonUtils {
}
}
public static Map<String, Object> parseMap(String text) {
if (StrUtil.isEmpty(text)) {
return null;
}
return parseObject(text, STRING_OBJECT_MAP);
}
public static List<Map<String, Object>> parseMapList(String text) {
if (StrUtil.isEmpty(text)) {
return null;
}
return parseObject(text, STRING_OBJECT_MAP_LIST);
}
public static Map<String, Object> toStringObjectMap(Object value) {
if (value == null) {
return null;
}
if (!(value instanceof Map<?, ?> map)) {
throw new ClassCastException("Expected Map but got " + value.getClass().getName());
}
Map<String, Object> result = new LinkedHashMap<>(map.size());
for (Map.Entry<?, ?> entry : map.entrySet()) {
result.put(String.class.cast(entry.getKey()), entry.getValue());
}
return result;
}
public static List<Map<String, Object>> toStringObjectMapList(Object value) {
if (value == null) {
return null;
}
List<?> list = List.class.cast(value);
List<Map<String, Object>> result = new ArrayList<>(list.size());
for (Object item : list) {
result.add(toStringObjectMap(item));
}
return result;
}
public static <T> List<T> toList(Object value, Class<T> elementType) {
if (value == null) {
return null;
}
List<?> list = List.class.cast(value);
List<T> result = new ArrayList<>(list.size());
for (Object item : list) {
result.add(elementType.cast(item));
}
return result;
}
}
@@ -75,11 +75,11 @@ public class SensitiveDataUtils {
Object value = jsonObject.get(key);
if (SENSITIVE_FIELDS.contains(key.toLowerCase()) && value instanceof String) {
result.put(key, maskMiddle((String) value));
result.set(key, maskMiddle((String) value));
} else if (value instanceof JSONObject) {
result.put(key, maskSensitiveFields((JSONObject) value));
result.set(key, maskSensitiveFields((JSONObject) value));
} else {
result.put(key, value);
result.set(key, value);
}
}
@@ -1,89 +0,0 @@
package xiaozhi.common.utils;
import cn.hutool.core.util.ReUtil;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import java.time.LocalDateTime;
import java.time.ZoneId;
import java.util.Date;
import java.util.List;
import java.util.Map;
import java.util.Set;
/**
* 通用工具类
*/
public class ToolUtil {
private static final Logger logger = LoggerFactory.getLogger(ToolUtil.class);
/**
* 对象是否不为空(新增)
*/
public static boolean isNotEmpty(Object o) {
return !isEmpty(o);
}
/**
* 对象是否为空
*/
public static boolean isEmpty(Object o) {
if (o == null) {
return true;
}
if (o instanceof String) {
if (o.toString().trim().equals("")) {
return true;
}
} else if (o instanceof List) {
if (((List) o).size() == 0) {
return true;
}
} else if (o instanceof Map) {
if (((Map) o).size() == 0) {
return true;
}
} else if (o instanceof Set) {
if (((Set) o).size() == 0) {
return true;
}
} else if (o instanceof Object[]) {
if (((Object[]) o).length == 0) {
return true;
}
} else if (o instanceof int[]) {
if (((int[]) o).length == 0) {
return true;
}
} else if (o instanceof long[]) {
if (((long[]) o).length == 0) {
return true;
}
}
return false;
}
/**
* 对象组中是否存在空对象
*/
public static boolean isOneEmpty(Object... os) {
for (Object o : os) {
if (isEmpty(o)) {
return true;
}
}
return false;
}
/**
* 对象组中是否全是空对象
*/
public static boolean isAllEmpty(Object... os) {
for (Object o : os) {
if (!isEmpty(o)) {
return false;
}
}
return true;
}
}
@@ -22,10 +22,10 @@ public class SqlFilter {
return null;
}
// 去掉'|"|;|\字符
str = StringUtils.replace(str, "'", "");
str = StringUtils.replace(str, "\"", "");
str = StringUtils.replace(str, ";", "");
str = StringUtils.replace(str, "\\", "");
str = str.replace("'", "");
str = str.replace("\"", "");
str = str.replace(";", "");
str = str.replace("\\", "");
// 转换成小写
str = str.toLowerCase();
@@ -34,6 +34,7 @@ import xiaozhi.common.page.PageData;
import xiaozhi.common.redis.RedisKeys;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.common.user.UserDetail;
import xiaozhi.common.utils.JsonUtils;
import xiaozhi.common.utils.Result;
import xiaozhi.common.utils.ResultUtils;
import xiaozhi.modules.agent.dto.AgentChatHistoryDTO;
@@ -340,8 +341,8 @@ public class AgentController {
@RequiresPermissions("sys:role:normal")
public Result<Void> saveAgentTags(@PathVariable String id, @RequestBody Map<String, Object> params) {
requireAgentPermission(id);
List<String> tagIds = (List<String>) params.get("tagIds");
List<String> tagNames = (List<String>) params.get("tagNames");
List<String> tagIds = JsonUtils.toList(params.get("tagIds"), String.class);
List<String> tagNames = JsonUtils.toList(params.get("tagNames"), String.class);
AgentUpdateDTO dto = new AgentUpdateDTO();
dto.setTagIds(tagIds);
dto.setTagNames(tagNames);
@@ -39,7 +39,7 @@ public class AgentDTO {
private String systemPrompt;
@Schema(description = "总结记忆", example = "构建可生长的动态记忆网络,在有限空间内保留关键信息的同时,智能维护信息演变轨迹\n" +
"根据对话记录,总结user的重要信息,以便在未来的对话中提供更个性化的服务", required = false)
"根据对话记录,总结user的重要信息,以便在未来的对话中提供更个性化的服务", requiredMode = Schema.RequiredMode.NOT_REQUIRED)
private String summaryMemory;
@Schema(description = "最后连接时间", example = "2024-03-20 10:00:00")
@@ -14,6 +14,6 @@ public class AgentMemoryDTO implements Serializable {
private static final long serialVersionUID = 1L;
@Schema(description = "总结记忆", example = "构建可生长的动态记忆网络,在有限空间内保留关键信息的同时,智能维护信息演变轨迹\n" +
"根据对话记录,总结user的重要信息,以便在未来的对话中提供更个性化的服务", required = false)
"根据对话记录,总结user的重要信息,以便在未来的对话中提供更个性化的服务", requiredMode = Schema.RequiredMode.NOT_REQUIRED)
private String summaryMemory;
}
@@ -39,10 +39,10 @@ public class AgentUpdateDTO implements Serializable {
@Schema(description = "小模型标识", example = "slm_model_02", nullable = true)
private String slmModelId;
@Schema(description = "VLLM模型标识", example = "vllm_model_02", required = false)
@Schema(description = "VLLM模型标识", example = "vllm_model_02", requiredMode = Schema.RequiredMode.NOT_REQUIRED)
private String vllmModelId;
@Schema(description = "语音合成模型标识", example = "tts_model_02", required = false)
@Schema(description = "语音合成模型标识", example = "tts_model_02", requiredMode = Schema.RequiredMode.NOT_REQUIRED)
private String ttsModelId;
@Schema(description = "音色标识", example = "voice_02", nullable = true)
@@ -74,7 +74,7 @@ public class AgentEntity {
private String systemPrompt;
@Schema(description = "总结记忆", example = "构建可生长的动态记忆网络,在有限空间内保留关键信息的同时,智能维护信息演变轨迹\n" +
"根据对话记录,总结user的重要信息,以便在未来的对话中提供更个性化的服务", required = false)
"根据对话记录,总结user的重要信息,以便在未来的对话中提供更个性化的服务", requiredMode = Schema.RequiredMode.NOT_REQUIRED)
private String summaryMemory;
@Schema(description = "语言编码")
@@ -5,6 +5,7 @@ import java.util.List;
import java.util.Map;
import java.util.stream.Collectors;
import cn.hutool.core.collection.CollUtil;
import cn.hutool.core.collection.ListUtil;
import lombok.RequiredArgsConstructor;
import org.springframework.stereotype.Service;
@@ -20,7 +21,6 @@ import xiaozhi.common.constant.Constant;
import xiaozhi.common.page.PageData;
import xiaozhi.common.utils.ConvertUtils;
import xiaozhi.common.utils.JsonUtils;
import xiaozhi.common.utils.ToolUtil;
import xiaozhi.modules.agent.Enums.AgentChatHistoryType;
import xiaozhi.modules.agent.dao.AiAgentChatHistoryDao;
import xiaozhi.modules.agent.dto.AgentChatHistoryDTO;
@@ -107,7 +107,7 @@ public class AgentChatHistoryServiceImpl extends CrudRepository<AiAgentChatHisto
if (deleteAudio) {
// 分批删除音频,避免超时
List<String> audioIds = baseMapper.getAudioIdsByAgentId(agentId);
if (ToolUtil.isNotEmpty(audioIds)) {
if (CollUtil.isNotEmpty(audioIds)) {
// 每批删除1000条
List<List<String>> batch = ListUtil.split(audioIds, 1000);
batch.forEach(dataList -> {
@@ -167,7 +167,7 @@ public class AgentChatHistoryServiceImpl extends CrudRepository<AiAgentChatHisto
// 尝试解析为 JSON
try {
Map<String, Object> jsonMap = JsonUtils.parseObject(content, Map.class);
Map<String, Object> jsonMap = JsonUtils.parseMap(content);
if (jsonMap != null && jsonMap.containsKey("content")) {
Object contentObj = jsonMap.get("content");
return contentObj != null ? contentObj.toString() : content;
@@ -12,11 +12,11 @@ import java.util.stream.Collectors;
import org.apache.commons.lang3.StringUtils;
import org.springframework.stereotype.Service;
import cn.hutool.crypto.digest.DigestUtil;
import lombok.AllArgsConstructor;
import lombok.extern.slf4j.Slf4j;
import xiaozhi.common.constant.Constant;
import xiaozhi.common.utils.AESUtils;
import xiaozhi.common.utils.HashEncryptionUtil;
import xiaozhi.common.utils.JsonUtils;
import xiaozhi.modules.agent.Enums.XiaoZhiMcpJsonRpcJson;
import xiaozhi.modules.agent.service.AgentMcpAccessPointService;
@@ -76,7 +76,7 @@ public class AgentMcpAccessPointServiceImpl implements AgentMcpAccessPointServic
// 等待初始化响应 (id=1) - 移除固定延迟,改为响应驱动
List<String> initResponses = client.listenerWithoutClose(response -> {
try {
Map<String, Object> jsonMap = JsonUtils.parseObject(response, Map.class);
Map<String, Object> jsonMap = JsonUtils.parseMap(response);
if (jsonMap != null && Integer.valueOf(1).equals(jsonMap.get("id"))) {
// 检查是否有result字段,表示初始化成功
return jsonMap.containsKey("result") && !jsonMap.containsKey("error");
@@ -92,7 +92,7 @@ public class AgentMcpAccessPointServiceImpl implements AgentMcpAccessPointServic
boolean initSucceeded = false;
for (String response : initResponses) {
try {
Map<String, Object> jsonMap = JsonUtils.parseObject(response, Map.class);
Map<String, Object> jsonMap = JsonUtils.parseMap(response);
if (jsonMap != null && Integer.valueOf(1).equals(jsonMap.get("id"))) {
if (jsonMap.containsKey("result")) {
log.info("MCP初始化成功,智能体ID: {}", id);
@@ -123,7 +123,7 @@ public class AgentMcpAccessPointServiceImpl implements AgentMcpAccessPointServic
// 等待工具列表响应 (id=2)
List<String> toolsResponses = client.listener(response -> {
try {
Map<String, Object> jsonMap = JsonUtils.parseObject(response, Map.class);
Map<String, Object> jsonMap = JsonUtils.parseMap(response);
return jsonMap != null && Integer.valueOf(2).equals(jsonMap.get("id"));
} catch (Exception e) {
log.warn("解析工具列表响应失败: {}", response, e);
@@ -134,18 +134,18 @@ public class AgentMcpAccessPointServiceImpl implements AgentMcpAccessPointServic
// 处理工具列表响应
for (String response : toolsResponses) {
try {
Map<String, Object> jsonMap = JsonUtils.parseObject(response, Map.class);
Map<String, Object> jsonMap = JsonUtils.parseMap(response);
if (jsonMap != null && Integer.valueOf(2).equals(jsonMap.get("id"))) {
// 检查是否有result字段
Object resultObj = jsonMap.get("result");
if (resultObj instanceof Map) {
Map<String, Object> resultMap = (Map<String, Object>) resultObj;
if (resultObj instanceof Map<?, ?>) {
Map<String, Object> resultMap = JsonUtils.toStringObjectMap(resultObj);
Object toolsObj = resultMap.get("tools");
if (toolsObj instanceof List) {
List<Map<String, Object>> toolsList = (List<Map<String, Object>>) toolsObj;
if (toolsObj instanceof List<?>) {
List<Map<String, Object>> toolsList = JsonUtils.toStringObjectMapList(toolsObj);
// 提取工具名称列表
List<String> result = toolsList.stream()
.map(tool -> (String) tool.get("name"))
.map(tool -> String.class.cast(tool.get("name")))
.filter(name -> name != null)
.sorted()
.collect(Collectors.toList());
@@ -226,7 +226,7 @@ public class AgentMcpAccessPointServiceImpl implements AgentMcpAccessPointServic
*/
private static String encryptToken(String agentId, String key) {
// 使用md5对智能体id进行加密
String md5 = HashEncryptionUtil.Md5hexDigest(agentId);
String md5 = DigestUtil.md5Hex(agentId);
// aes需要加密文本
String json = "{\"agentId\": \"%s\"}".formatted(md5);
// 加密后成token值
@@ -18,6 +18,7 @@ import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.core.metadata.IPage;
import com.baomidou.mybatisplus.extension.repository.IRepository;
import cn.hutool.core.collection.CollUtil;
import lombok.AllArgsConstructor;
import xiaozhi.common.constant.Constant;
import xiaozhi.common.exception.ErrorCode;
@@ -29,7 +30,6 @@ import xiaozhi.common.service.impl.BaseServiceImpl;
import xiaozhi.common.user.UserDetail;
import xiaozhi.common.utils.ConvertUtils;
import xiaozhi.common.utils.JsonUtils;
import xiaozhi.common.utils.ToolUtil;
import xiaozhi.modules.agent.dao.AgentDao;
import xiaozhi.modules.agent.dao.AgentTagDao;
import xiaozhi.modules.agent.dto.AgentCreateDTO;
@@ -243,13 +243,13 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
.map(DeviceEntity::getAgentId)
.distinct()
.collect(Collectors.toList());
if (ToolUtil.isNotEmpty(agentIds)) {
if (CollUtil.isNotEmpty(agentIds)) {
w.or().in("id", agentIds);
}
// 按标签名搜索
List<String> tagAgentIds = agentTagService.getAgentIdsByTagName(keyword);
if (ToolUtil.isNotEmpty(tagAgentIds)) {
if (CollUtil.isNotEmpty(tagAgentIds)) {
w.or().in("id", tagAgentIds);
}
});
@@ -291,7 +291,7 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
// 获取标签列表
List<AgentTagEntity> tags = agentTagDao.selectByAgentId(agent.getId());
if (ToolUtil.isNotEmpty(tags)) {
if (CollUtil.isNotEmpty(tags)) {
dto.setTags(tags.stream().map(this::convertTagToDTO).collect(Collectors.toList()));
}
@@ -677,7 +677,7 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
mapping.setPluginId(pluginId);
Map<String, Object> paramInfo = new HashMap<>();
List<Map<String, Object>> fields = JsonUtils.parseObject(provider.getFields(), List.class);
List<Map<String, Object>> fields = JsonUtils.parseMapList(provider.getFields());
if (fields != null) {
for (Map<String, Object> field : fields) {
paramInfo.put((String) field.get("key"), field.get("default"));
@@ -133,7 +133,7 @@ public class AgentTagServiceImpl extends BaseServiceImpl<AgentTagDao, AgentTagEn
}
if (tagIds != null && !tagIds.isEmpty()) {
List<AgentTagEntity> tagIdEntities = baseDao.selectBatchIds(tagIds);
List<AgentTagEntity> tagIdEntities = baseDao.selectByIds(tagIds);
for (AgentTagEntity tag : tagIdEntities) {
if (tag != null && (currentTagNames.contains(tag.getTagName()) ||
newTagNames.contains(tag.getTagName()))) {
@@ -10,7 +10,7 @@ public interface ConfigService {
* @param isCache 是否缓存
* @return 配置信息
*/
Object getConfig(Boolean isCache);
Map<String, Object> getConfig(Boolean isCache);
/**
* 获取智能体模型配置
@@ -65,12 +65,12 @@ public class ConfigServiceImpl implements ConfigService {
private final CorrectWordFileService correctWordFileService;
@Override
public Object getConfig(Boolean isCache) {
public Map<String, Object> getConfig(Boolean isCache) {
if (isCache) {
// 先从Redis获取配置
Object cachedConfig = redisUtils.get(RedisKeys.getServerConfigKey());
if (cachedConfig != null) {
return cachedConfig;
return JsonUtils.toStringObjectMap(cachedConfig);
}
}
@@ -123,7 +123,7 @@ public class ConfigServiceImpl implements ConfigService {
if (isAdminRequest != null && "true".equals(isAdminRequest)) {
// 管理控制台请求,返回getConfig的结果
redisUtils.delete(redisKey); // 使用后清理
return (Map<String, Object>) getConfig(true);
return getConfig(true);
}
// 根据MAC地址查找设备
DeviceEntity device = deviceService.getDeviceByMacAddress(macAddress);
@@ -277,10 +277,10 @@ public class ConfigServiceImpl implements ConfigService {
// 遍历除最后一个key之外的所有key
for (int i = 0; i < keys.length - 1; i++) {
String key = keys[i];
if (!current.containsKey(key)) {
current.put(key, new HashMap<String, Object>());
}
current = (Map<String, Object>) current.get(key);
Object nestedConfig = current.computeIfAbsent(key, ignored -> new HashMap<String, Object>());
Map<String, Object> nestedMap = JsonUtils.toStringObjectMap(nestedConfig);
current.put(key, nestedMap);
current = nestedMap;
}
// 处理最后一个key
@@ -125,7 +125,9 @@ public class DeviceController {
return new Result<Void>().error("设备不存在");
}
BeanUtils.copyProperties(deviceUpdateDTO, entity);
deviceService.updateById(entity);
if (!deviceService.updateById(entity)) {
return new Result<Void>().error(ErrorCode.UPDATE_DATA_FAILED);
}
return new Result<Void>();
}
@@ -5,8 +5,6 @@ import java.io.IOException;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.security.MessageDigest;
import java.security.NoSuchAlgorithmException;
import java.util.List;
import java.util.Map;
import java.util.Optional;
@@ -31,6 +29,7 @@ import org.springframework.web.bind.annotation.RequestParam;
import org.springframework.web.bind.annotation.RestController;
import org.springframework.web.multipart.MultipartFile;
import cn.hutool.crypto.digest.DigestUtil;
import io.swagger.v3.oas.annotations.Operation;
import io.swagger.v3.oas.annotations.Parameter;
import io.swagger.v3.oas.annotations.Parameters;
@@ -288,7 +287,7 @@ public class OTAMagController {
// 返回文件路径
return new Result<String>().ok(filePath.toString());
} catch (IOException | NoSuchAlgorithmException e) {
} catch (IOException e) {
return new Result<String>().error("文件上传失败:" + e.getMessage());
}
}
@@ -329,13 +328,7 @@ public class OTAMagController {
return result;
}
private String calculateMD5(MultipartFile file) throws IOException, NoSuchAlgorithmException {
MessageDigest md = MessageDigest.getInstance("MD5");
byte[] digest = md.digest(file.getBytes());
StringBuilder sb = new StringBuilder();
for (byte b : digest) {
sb.append(String.format("%02x", b));
}
return sb.toString();
private String calculateMD5(MultipartFile file) throws IOException {
return DigestUtil.md5Hex(file.getBytes());
}
}
@@ -12,6 +12,11 @@ import xiaozhi.modules.device.entity.DeviceAddressBookEntity;
@Mapper
public interface DeviceAddressBookDao extends BaseMapper<DeviceAddressBookEntity> {
/**
* 新增设备通讯录记录
*/
int insertAddressBook(DeviceAddressBookEntity entity);
/**
* 获取设备通讯录列表
*/
@@ -167,7 +167,7 @@ public class DeviceAddressBookServiceImpl implements DeviceAddressBookService {
alias = generateUniqueAlias(macAddress, targetMac, alias);
entity.setAlias(alias);
entity.setHasPermission(hasPermission);
deviceAddressBookDao.insert(entity);
deviceAddressBookDao.insertAddressBook(entity);
} else {
if (alias != null) {
updateAlias(macAddress, targetMac, alias);
@@ -31,6 +31,7 @@ import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.core.conditions.update.UpdateWrapper;
import com.baomidou.mybatisplus.core.metadata.IPage;
import cn.hutool.core.collection.CollUtil;
import cn.hutool.core.map.MapUtil;
import cn.hutool.core.util.RandomUtil;
import cn.hutool.core.util.StrUtil;
@@ -50,7 +51,7 @@ import xiaozhi.common.service.impl.BaseServiceImpl;
import xiaozhi.common.user.UserDetail;
import xiaozhi.common.utils.ConvertUtils;
import xiaozhi.common.utils.DateUtils;
import xiaozhi.common.utils.ToolUtil;
import xiaozhi.common.utils.JsonUtils;
import xiaozhi.modules.device.dao.DeviceDao;
import xiaozhi.modules.device.dto.DeviceManualAddDTO;
import xiaozhi.modules.device.dto.DevicePageUserDTO;
@@ -102,15 +103,15 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
throw new RenException(ErrorCode.ACTIVATION_CODE_EMPTY);
}
String deviceKey = RedisKeys.getOtaActivationCode(activationCode);
Object cacheDeviceId = redisUtils.get(deviceKey);
if (ToolUtil.isEmpty(cacheDeviceId)) {
String cacheDeviceId = (String) redisUtils.get(deviceKey);
if (StringUtils.isBlank(cacheDeviceId)) {
throw new RenException(ErrorCode.ACTIVATION_CODE_ERROR);
}
String deviceId = (String) cacheDeviceId;
String deviceId = cacheDeviceId;
String safeDeviceId = deviceId.replace(":", "_").toLowerCase();
String cacheDeviceKey = RedisKeys.getOtaDeviceActivationInfo(safeDeviceId);
Map<String, Object> cacheMap = (Map<String, Object>) redisUtils.get(cacheDeviceKey);
if (ToolUtil.isEmpty(cacheMap)) {
Map<String, Object> cacheMap = JsonUtils.toStringObjectMap(redisUtils.get(cacheDeviceKey));
if (MapUtil.isEmpty(cacheMap)) {
throw new RenException(ErrorCode.ACTIVATION_CODE_ERROR);
}
String cachedCode = (String) cacheMap.get("activation_code");
@@ -180,7 +181,7 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
.builder(new HashMap<String, Set<String>>())
.put("clientIds", deviceIds).build();
if (ToolUtil.isNotEmpty(deviceIds)) {
if (CollUtil.isNotEmpty(deviceIds)) {
return postToMqttGateway(url, params);
}
// 返回响应
@@ -201,8 +202,8 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
firmware.setUrl(Constant.INVALID_FIRMWARE_URL);
response.setFirmware(firmware);
} else {
// 只有在设备已绑定且autoUpdate不为0的情况下才返回固件升级信息
if (deviceById.getAutoUpdate() != 0) {
// 只有在设备已绑定且明确开启自动升级时才返回固件升级信息
if (Integer.valueOf(1).equals(deviceById.getAutoUpdate())) {
String type = deviceReport.getBoard() == null ? null : deviceReport.getBoard().getType();
DeviceReportRespDTO.Firmware firmware = buildFirmwareInfo(type,
deviceReport.getApplication() == null ? null : deviceReport.getApplication().getVersion());
@@ -298,6 +299,7 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
UserShowDeviceListVO vo = ConvertUtils.sourceToTarget(device, UserShowDeviceListVO.class);
vo.setDeviceType(device.getBoard());
vo.setBoard(device.getBoard());
vo.setAutoUpdate(device.getAutoUpdate());
vo.setCreateDateTimestamp(toTimestamp(device.getCreateDate()));
vo.setLastConnectedAtTimestamp(toTimestamp(device.getLastConnectedAt()));
return vo;
@@ -411,7 +413,7 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
public String geCodeByDeviceId(String deviceId) {
String dataKey = getDeviceCacheKey(deviceId);
Map<String, Object> cacheMap = (Map<String, Object>) redisUtils.get(dataKey);
Map<String, Object> cacheMap = JsonUtils.toStringObjectMap(redisUtils.get(dataKey));
if (cacheMap != null && cacheMap.containsKey("activation_code")) {
String cachedCode = (String) cacheMap.get("activation_code");
return cachedCode;
@@ -56,7 +56,7 @@ public class OtaServiceImpl extends BaseServiceImpl<OtaDao, OtaEntity> implement
@Override
public void delete(String[] ids) {
baseDao.deleteBatchIds(Arrays.asList(ids));
baseDao.deleteByIds(Arrays.asList(ids));
}
@Override
@@ -32,8 +32,8 @@ public class UserShowDeviceListVO {
@Schema(description = "设备别名")
private String alias;
@Schema(description = "开启OTA")
private Integer otaUpgrade;
@Schema(description = "自动更新开关(0关闭/1开启)")
private Integer autoUpdate;
@Schema(description = "最近对话时间")
private String recentChatTime;
@@ -15,6 +15,7 @@ import org.springframework.web.bind.annotation.RequestMapping;
import org.springframework.web.bind.annotation.RequestParam;
import org.springframework.web.bind.annotation.RestController;
import cn.hutool.core.collection.CollUtil;
import io.swagger.v3.oas.annotations.Operation;
import io.swagger.v3.oas.annotations.Parameter;
import io.swagger.v3.oas.annotations.tags.Tag;
@@ -23,7 +24,6 @@ import xiaozhi.common.exception.ErrorCode;
import xiaozhi.common.exception.RenException;
import xiaozhi.common.page.PageData;
import xiaozhi.common.utils.Result;
import xiaozhi.common.utils.ToolUtil;
import xiaozhi.modules.knowledge.dto.KnowledgeBaseDTO;
import xiaozhi.modules.knowledge.service.KnowledgeBaseService;
import xiaozhi.modules.knowledge.service.KnowledgeManagerService;
@@ -140,7 +140,7 @@ public class KnowledgeBaseController {
List<String> idList = Arrays.asList(ids.split(","));
List<KnowledgeBaseDTO> knowledgeBaseDTOs = Optional.ofNullable(knowledgeBaseService.getByDatasetIdList(idList))
.orElseGet(ArrayList::new);
if (ToolUtil.isNotEmpty(knowledgeBaseDTOs)) {
if (CollUtil.isNotEmpty(knowledgeBaseDTOs)) {
knowledgeBaseDTOs.forEach(item -> {
// 检查权限:用户只能删除自己创建的知识库
if (item.getCreator() == null || !item.getCreator().equals(currentUserId)) {
@@ -149,8 +149,8 @@ public class RAGFlowAdapter extends KnowledgeBaseAdapter {
log.info("=== [RAGFlow] 获取文档列表: datasetId={} ===", datasetId);
// 使用 Jackson 将 DTO 转为 Map 作为查询参数
@SuppressWarnings("unchecked")
Map<String, Object> params = objectMapper.convertValue(req, Map.class);
Map<String, Object> params = objectMapper.convertValue(req, new TypeReference<Map<String, Object>>() {
});
Map<String, Object> response = getClient().get("/api/v1/datasets/" + datasetId + "/documents", params);
@@ -174,13 +174,12 @@ public class RAGFlowAdapter extends KnowledgeBaseAdapter {
.pageSize(1)
.build();
@SuppressWarnings("unchecked")
Map<String, Object> params = objectMapper.convertValue(req, Map.class);
Map<String, Object> params = objectMapper.convertValue(req, new TypeReference<Map<String, Object>>() {
});
Map<String, Object> response = getClient().get("/api/v1/datasets/" + datasetId + "/documents", params);
Object dataObj = response.get("data");
if (dataObj instanceof Map) {
Map<String, Object> dataMap = (Map<String, Object>) dataObj;
if (dataObj instanceof Map<?, ?> dataMap) {
List<?> documents = (List<?>) dataMap.get("docs");
if (documents != null && !documents.isEmpty()) {
return objectMapper.convertValue(documents.get(0), DocumentDTO.InfoVO.class);
@@ -567,8 +566,8 @@ public class RAGFlowAdapter extends KnowledgeBaseAdapter {
return new PageData<>(new ArrayList<>(), 0);
}
Map<String, Object> dataMap = (Map<String, Object>) dataObj;
List<Map<String, Object>> documents = (List<Map<String, Object>>) dataMap.get("docs");
Map<?, ?> dataMap = Map.class.cast(dataObj);
List<?> documents = List.class.cast(dataMap.get("docs"));
if (documents == null || documents.isEmpty()) {
// RAGFlow 明确返回了空文档列表,这是合法的"真空"
return new PageData<>(new ArrayList<>(), 0);
@@ -679,7 +678,9 @@ public class RAGFlowAdapter extends KnowledgeBaseAdapter {
dto.setChunkMethod(info.getChunkMethod().name().toLowerCase());
}
if (info.getParserConfig() != null) {
dto.setParserConfig(objectMapper.convertValue(info.getParserConfig(), Map.class));
dto.setParserConfig(objectMapper.convertValue(info.getParserConfig(),
new TypeReference<Map<String, Object>>() {
}));
}
return dto;
@@ -116,13 +116,11 @@ public class OpenAIStyleLLMServiceImpl implements LLMService {
Map<String, Object> requestBody = new HashMap<>();
requestBody.put("model", model != null ? model : "gpt-3.5-turbo");
Map<String, Object>[] messages = new Map[1];
Map<String, Object> message = new HashMap<>();
message.put("role", "user");
message.put("content", prompt);
messages[0] = message;
requestBody.put("messages", messages);
requestBody.put("messages", List.of(message));
requestBody.put("temperature", temperature != null ? temperature : 0.7);
requestBody.put("max_tokens", maxTokens != null ? maxTokens : 2000);
@@ -212,13 +210,11 @@ public class OpenAIStyleLLMServiceImpl implements LLMService {
Map<String, Object> requestBody = new HashMap<>();
requestBody.put("model", model != null ? model : "gpt-3.5-turbo");
Map<String, Object>[] messages = new Map[1];
Map<String, Object> message = new HashMap<>();
message.put("role", "user");
message.put("content", prompt);
messages[0] = message;
requestBody.put("messages", messages);
requestBody.put("messages", List.of(message));
requestBody.put("temperature", 0.2);
requestBody.put("max_tokens", 2000);
@@ -368,13 +364,11 @@ public class OpenAIStyleLLMServiceImpl implements LLMService {
Map<String, Object> requestBody = new HashMap<>();
requestBody.put("model", model != null ? model : "gpt-3.5-turbo");
Map<String, Object>[] messages = new Map[1];
Map<String, Object> message = new HashMap<>();
message.put("role", "user");
message.put("content", prompt);
messages[0] = message;
requestBody.put("messages", messages);
requestBody.put("messages", List.of(message));
requestBody.put("temperature", 0.3);
requestBody.put("max_tokens", 50);
@@ -362,14 +362,14 @@ public class ModelConfigServiceImpl extends BaseServiceImpl<ModelConfigDao, Mode
if (SensitiveDataUtils.isSensitiveField(key)) {
if (value instanceof String && !SensitiveDataUtils.isMaskedValue((String) value)) {
updatedJson.put(key, value);
updatedJson.set(key, value);
}
} else if (value instanceof JSONObject) {
// 递归处理嵌套JSON
mergeJson(updatedJson, key, (JSONObject) value);
} else {
// 非敏感字段直接更新
updatedJson.put(key, value);
updatedJson.set(key, value);
}
}
@@ -405,7 +405,7 @@ public class ModelConfigServiceImpl extends BaseServiceImpl<ModelConfigDao, Mode
// 如果 original 中不存在 key,创建一个新的 JSON 对象
if (!original.containsKey(key)) {
original.put(key, new JSONObject());
original.set(key, new JSONObject());
}
// 获取 original 中的子对象
@@ -420,7 +420,7 @@ public class ModelConfigServiceImpl extends BaseServiceImpl<ModelConfigDao, Mode
log.warn("mergeJson: key '{}' 的值不是 JSONObject 类型 (实际类型:{}),将创建新对象",
key, originalValue != null ? originalValue.getClass().getSimpleName() : "null");
originalChild = new JSONObject();
original.put(key, originalChild);
original.set(key, originalChild);
}
for (String childKey : updated.keySet()) {
@@ -430,7 +430,7 @@ public class ModelConfigServiceImpl extends BaseServiceImpl<ModelConfigDao, Mode
} else {
if (!SensitiveDataUtils.isSensitiveField(childKey) ||
(childValue instanceof String && !isMaskedValue((String) childValue))) {
originalChild.put(childKey, childValue);
originalChild.set(childKey, childValue);
}
}
}
@@ -164,7 +164,7 @@ public class ModelProviderServiceImpl extends BaseServiceImpl<ModelProviderDao,
@Override
public void delete(List<String> ids) {
if (modelProviderDao.deleteBatchIds(ids) == 0) {
if (modelProviderDao.deleteByIds(ids) == 0) {
throw new RenException(ErrorCode.DELETE_DATA_FAILED);
}
}
@@ -4,16 +4,19 @@ import java.util.HashMap;
import java.util.LinkedHashMap;
import java.util.Map;
import org.apache.shiro.mgt.SecurityManager;
import org.apache.shiro.session.mgt.SessionManager;
import org.apache.shiro.spring.LifecycleBeanPostProcessor;
import org.apache.shiro.spring.security.interceptor.AuthorizationAttributeSourceAdvisor;
import org.apache.shiro.spring.web.ShiroFilterFactoryBean;
import org.apache.shiro.web.config.ShiroFilterConfiguration;
import org.apache.shiro.web.mgt.DefaultWebSecurityManager;
import org.apache.shiro.web.mgt.WebSecurityManager;
import org.apache.shiro.web.session.mgt.DefaultWebSessionManager;
import org.springframework.beans.factory.config.BeanDefinition;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.context.annotation.Lazy;
import org.springframework.context.annotation.Role;
import jakarta.servlet.Filter;
import xiaozhi.modules.security.oauth2.Oauth2Filter;
@@ -39,7 +42,7 @@ public class ShiroConfig {
}
@Bean("securityManager")
public SecurityManager securityManager(Oauth2Realm oAuth2Realm, SessionManager sessionManager) {
public WebSecurityManager securityManager(Oauth2Realm oAuth2Realm, SessionManager sessionManager) {
DefaultWebSecurityManager securityManager = new DefaultWebSecurityManager();
securityManager.setRealm(oAuth2Realm);
securityManager.setSessionManager(sessionManager);
@@ -48,7 +51,8 @@ public class ShiroConfig {
}
@Bean("shiroFilter")
public ShiroFilterFactoryBean shirFilter(SecurityManager securityManager, SysParamsService sysParamsService) {
public static ShiroFilterFactoryBean shirFilter(@Lazy WebSecurityManager securityManager,
@Lazy SysParamsService sysParamsService) {
ShiroFilterConfiguration config = new ShiroFilterConfiguration();
config.setFilterOncePerRequest(true);
@@ -101,12 +105,14 @@ public class ShiroConfig {
}
@Bean("lifecycleBeanPostProcessor")
public LifecycleBeanPostProcessor lifecycleBeanPostProcessor() {
public static LifecycleBeanPostProcessor lifecycleBeanPostProcessor() {
return new LifecycleBeanPostProcessor();
}
@Bean
public AuthorizationAttributeSourceAdvisor authorizationAttributeSourceAdvisor(SecurityManager securityManager) {
@Role(BeanDefinition.ROLE_INFRASTRUCTURE)
public static AuthorizationAttributeSourceAdvisor authorizationAttributeSourceAdvisor(
@Lazy WebSecurityManager securityManager) {
AuthorizationAttributeSourceAdvisor advisor = new AuthorizationAttributeSourceAdvisor();
advisor.setSecurityManager(securityManager);
return advisor;
@@ -1,8 +1,7 @@
package xiaozhi.modules.security.service.impl;
import java.io.IOException;
import java.util.Random;
import java.util.concurrent.TimeUnit;
import java.time.Duration;
import org.apache.commons.lang3.StringUtils;
import org.springframework.beans.factory.annotation.Value;
@@ -13,6 +12,7 @@ import com.google.common.cache.CacheBuilder;
import com.wf.captcha.SpecCaptcha;
import com.wf.captcha.base.Captcha;
import cn.hutool.core.util.RandomUtil;
import jakarta.annotation.Resource;
import jakarta.servlet.http.HttpServletResponse;
import xiaozhi.common.constant.Constant;
@@ -41,7 +41,7 @@ public class CaptchaServiceImpl implements CaptchaService {
* Local Cache 5分钟过期
*/
Cache<String, String> localCache = CacheBuilder.newBuilder().maximumSize(1000)
.expireAfterAccess(5, TimeUnit.MINUTES).build();
.expireAfterAccess(Duration.ofMinutes(5)).build();
@Override
public void create(HttpServletResponse response, String uuid) throws IOException {
@@ -113,7 +113,7 @@ public class CaptchaServiceImpl implements CaptchaService {
}
String key = RedisKeys.getSMSValidateCodeKey(phone);
String validateCodes = generateValidateCode(6);
String validateCodes = RandomUtil.randomNumbers(6);
// 设置验证码
setCache(key, validateCodes);
@@ -135,22 +135,6 @@ public class CaptchaServiceImpl implements CaptchaService {
return validate(key, code, delete);
}
/**
* 生成指定数量的随机数验证码
*
* @param length 数量
* @return 随机码
*/
private String generateValidateCode(Integer length) {
String chars = "0123456789"; // 字符范围可以自定义:数字
Random random = new Random();
StringBuilder code = new StringBuilder();
for (int i = 0; i < length; i++) {
code.append(chars.charAt(random.nextInt(chars.length())));
}
return code.toString();
}
private void setCache(String key, String value) {
if (open) {
key = RedisKeys.getCaptchaKey(key);
@@ -31,7 +31,7 @@ public class SysUserDTO implements Serializable {
@NotNull(message = "{id.require}", groups = UpdateGroup.class)
private Long id;
@Schema(description = "用户名", required = true)
@Schema(description = "用户名", requiredMode = Schema.RequiredMode.REQUIRED)
@NotBlank(message = "{sysuser.username.require}", groups = DefaultGroup.class)
private String username;
@@ -40,14 +40,14 @@ public class SysUserDTO implements Serializable {
@NotBlank(message = "{sysuser.password.require}", groups = AddGroup.class)
private String password;
@Schema(description = "姓名", required = true)
@Schema(description = "姓名", requiredMode = Schema.RequiredMode.REQUIRED)
@NotBlank(message = "{sysuser.realname.require}", groups = DefaultGroup.class)
private String realName;
@Schema(description = "头像")
private String headUrl;
@Schema(description = "性别 0:男 1:女 2:保密", required = true)
@Schema(description = "性别 0:男 1:女 2:保密", requiredMode = Schema.RequiredMode.REQUIRED)
@Range(min = 0, max = 2, message = "{sysuser.gender.range}", groups = DefaultGroup.class)
private Integer gender;
@@ -58,11 +58,11 @@ public class SysUserDTO implements Serializable {
@Schema(description = "手机号")
private String mobile;
@Schema(description = "部门ID", required = true)
@Schema(description = "部门ID", requiredMode = Schema.RequiredMode.REQUIRED)
@NotNull(message = "{sysuser.deptId.require}", groups = DefaultGroup.class)
private Long deptId;
@Schema(description = "状态 0:停用 1:正常", required = true)
@Schema(description = "状态 0:停用 1:正常", requiredMode = Schema.RequiredMode.REQUIRED)
@Range(min = 0, max = 1, message = "{sysuser.status.range}", groups = DefaultGroup.class)
private Integer status;
@@ -12,6 +12,7 @@ import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.core.metadata.IPage;
import cn.hutool.core.collection.CollUtil;
import lombok.AllArgsConstructor;
import xiaozhi.common.exception.RenException;
import xiaozhi.common.exception.ErrorCode;
@@ -20,7 +21,7 @@ import xiaozhi.common.redis.RedisKeys;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.common.service.impl.BaseServiceImpl;
import xiaozhi.common.utils.ConvertUtils;
import xiaozhi.common.utils.ToolUtil;
import xiaozhi.common.utils.JsonUtils;
import xiaozhi.modules.sys.dao.SysDictDataDao;
import xiaozhi.modules.sys.dao.SysUserDao;
import xiaozhi.modules.sys.dto.SysDictDataDTO;
@@ -103,13 +104,13 @@ public class SysDictDataServiceImpl extends BaseServiceImpl<SysDictDataDao, SysD
@Transactional(rollbackFor = Exception.class)
public void delete(Long[] ids) {
List<Long> idList = Arrays.asList(ids);
if (ToolUtil.isNotEmpty(idList)) {
if (CollUtil.isNotEmpty(idList)) {
//批量删除redis字典
List<String> redisKeyList = new ArrayList<>();
//批量获取字典类型
List<String> dictTypeList = Optional.ofNullable(baseDao.getDictTypesByIdList(idList)).orElseGet(ArrayList::new);
dictTypeList.forEach(dictType -> redisKeyList.add(RedisKeys.getDictDataByTypeKey(dictType)));
if (ToolUtil.isNotEmpty(redisKeyList)) {
if (CollUtil.isNotEmpty(redisKeyList)) {
//清除缓存
redisUtils.delete(redisKeyList);
}
@@ -138,7 +139,7 @@ public class SysDictDataServiceImpl extends BaseServiceImpl<SysDictDataDao, SysD
// 设置更新者和创建者名称
if (!userIds.isEmpty()) {
List<SysUserEntity> sysUserEntities = sysUserDao.selectBatchIds(userIds);
List<SysUserEntity> sysUserEntities = sysUserDao.selectByIds(userIds);
// 把List转成MapMap<Long, String>
Map<Long, String> userNameMap = sysUserEntities.stream().collect(Collectors.toMap(SysUserEntity::getId,
SysUserEntity::getUsername, (existing, replacement) -> existing));
@@ -170,7 +171,7 @@ public class SysDictDataServiceImpl extends BaseServiceImpl<SysDictDataDao, SysD
// 先从Redis获取缓存
String key = RedisKeys.getDictDataByTypeKey(dictType);
List<SysDictDataItem> cachedData = (List<SysDictDataItem>) redisUtils.get(key);
List<SysDictDataItem> cachedData = JsonUtils.toList(redisUtils.get(key), SysDictDataItem.class);
if (cachedData != null) {
return cachedData;
}
@@ -128,7 +128,7 @@ public class SysDictTypeServiceImpl extends BaseServiceImpl<SysDictTypeDao, SysD
// 设置更新者和创建者名称
if (!userIds.isEmpty()) {
List<SysUserEntity> sysUserEntities = sysUserDao.selectBatchIds(userIds);
List<SysUserEntity> sysUserEntities = sysUserDao.selectByIds(userIds);
// 把List转成MapMap<Long, String>
Map<Long, String> userNameMap = sysUserEntities.stream().collect(Collectors.toMap(SysUserEntity::getId,
SysUserEntity::getUsername, (existing, replacement) -> existing));
@@ -287,10 +287,10 @@ public class SysParamsServiceImpl extends BaseServiceImpl<SysParamsDao, SysParam
try {
if (StringUtils.isNotBlank(currentConfig)) {
currentMap = JsonUtils.parseObject(currentConfig, Map.class);
currentMap = JsonUtils.parseMap(currentConfig);
}
if (StringUtils.isNotBlank(configJson)) {
newMap = JsonUtils.parseObject(configJson, Map.class);
newMap = JsonUtils.parseMap(configJson);
}
} catch (Exception e) {
throw new RenException(ErrorCode.PARAM_JSON_INVALID);
@@ -298,8 +298,8 @@ public class SysParamsServiceImpl extends BaseServiceImpl<SysParamsDao, SysParam
// 检查addressBook功能是否被关闭
if (currentMap != null && newMap != null) {
Map<String, Object> currentFeatures = (Map<String, Object>) currentMap.get("features");
Map<String, Object> newFeatures = (Map<String, Object>) newMap.get("features");
Map<?, ?> currentFeatures = Map.class.cast(currentMap.get("features"));
Map<?, ?> newFeatures = Map.class.cast(newMap.get("features"));
if (currentFeatures != null && newFeatures != null) {
Object currentAddressBookObj = currentFeatures.get("addressBook");
@@ -308,16 +308,14 @@ public class SysParamsServiceImpl extends BaseServiceImpl<SysParamsDao, SysParam
Boolean currentEnabled = false;
Boolean newEnabled = false;
if (currentAddressBookObj instanceof Map) {
Map<String, Object> currentAddressBook = (Map<String, Object>) currentAddressBookObj;
currentEnabled = currentAddressBook.get("enabled") != null
? (Boolean) currentAddressBook.get("enabled") : false;
if (currentAddressBookObj instanceof Map<?, ?> currentAddressBook) {
Object enabled = currentAddressBook.get("enabled");
currentEnabled = enabled != null ? Boolean.class.cast(enabled) : false;
}
if (newAddressBookObj instanceof Map) {
Map<String, Object> newAddressBook = (Map<String, Object>) newAddressBookObj;
newEnabled = newAddressBook.get("enabled") != null
? (Boolean) newAddressBook.get("enabled") : false;
if (newAddressBookObj instanceof Map<?, ?> newAddressBook) {
Object enabled = newAddressBook.get("enabled");
newEnabled = enabled != null ? Boolean.class.cast(enabled) : false;
}
// 如果之前是启用状态,现在被禁用,删除所有call_device插件
@@ -34,7 +34,7 @@ public class TimbreDataDTO {
@Schema(description = "排序")
@Min(value = 0, message = "{sort.number}")
private long sort;
private Long sort;
@Schema(description = "对应 TTS 模型主键")
@NotBlank(message = "{timbre.ttsModelId.require}")
@@ -3,6 +3,7 @@ package xiaozhi.modules.timbre.entity;
import java.util.Date;
import com.baomidou.mybatisplus.annotation.FieldFill;
import com.baomidou.mybatisplus.annotation.FieldStrategy;
import com.baomidou.mybatisplus.annotation.TableField;
import com.baomidou.mybatisplus.annotation.TableName;
@@ -41,7 +42,8 @@ public class TimbreEntity {
private String referenceText;
@Schema(description = "排序")
private long sort;
@TableField(updateStrategy = FieldStrategy.NOT_NULL)
private Long sort;
@Schema(description = "对应 TTS 模型主键")
private String ttsModelId;
@@ -99,6 +99,9 @@ public class TimbreServiceImpl extends BaseServiceImpl<TimbreDao, TimbreEntity>
@Transactional(rollbackFor = Exception.class)
public void save(TimbreDataDTO dto) {
isTtsModelId(dto.getTtsModelId());
if (dto.getSort() == null) {
dto.setSort(0L);
}
TimbreEntity timbreEntity = ConvertUtils.sourceToTarget(dto, TimbreEntity.class);
baseDao.insert(timbreEntity);
}
@@ -117,7 +120,7 @@ public class TimbreServiceImpl extends BaseServiceImpl<TimbreDao, TimbreEntity>
@Override
@Transactional(rollbackFor = Exception.class)
public void delete(String[] ids) {
baseDao.deleteBatchIds(Arrays.asList(ids));
baseDao.deleteByIds(Arrays.asList(ids));
}
@Override
@@ -32,7 +32,7 @@ public class TimbreDetailsVO implements Serializable {
private String referenceText;
@Schema(description = "排序")
private long sort;
private Long sort;
@Schema(description = "对应 TTS 模型主键")
private String ttsModelId;
@@ -17,6 +17,7 @@ import com.baomidou.mybatisplus.core.metadata.IPage;
import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.ObjectMapper;
import cn.hutool.core.collection.CollUtil;
import lombok.RequiredArgsConstructor;
import lombok.extern.slf4j.Slf4j;
import xiaozhi.common.constant.Constant;
@@ -26,7 +27,6 @@ import xiaozhi.common.page.PageData;
import xiaozhi.common.service.impl.BaseServiceImpl;
import xiaozhi.common.utils.ConvertUtils;
import xiaozhi.common.utils.DateUtils;
import xiaozhi.common.utils.ToolUtil;
import xiaozhi.modules.model.entity.ModelConfigEntity;
import xiaozhi.modules.model.service.ModelConfigService;
import xiaozhi.modules.sys.dao.SysUserDao;
@@ -120,14 +120,14 @@ public class VoiceCloneServiceImpl extends BaseServiceImpl<VoiceCloneDao, VoiceC
entity.setTrainStatus(0); // 默认训练中
batchInsertList.add(entity);
}
if (ToolUtil.isNotEmpty(batchInsertList)) {
if (CollUtil.isNotEmpty(batchInsertList)) {
insertBatch(batchInsertList);
}
}
@Override
public void delete(String[] ids) {
baseDao.deleteBatchIds(Arrays.asList(ids));
baseDao.deleteByIds(Arrays.asList(ids));
}
@Override
@@ -1,20 +1,8 @@
<?xml version="1.0" encoding="UTF-8"?>
<configuration>
<!-- 启用JansiConsoleAppender以确保控制台输出有颜色 -->
<conversionRule conversionWord="clr" converterClass="org.springframework.boot.logging.logback.ColorConverter" />
<!-- 确保日志目录存在 -->
<timestamp key="bySecond" datePattern="yyyyMMdd'T'HHmmss"/>
<!-- 定义日志文件存储位置 -->
<property name="LOG_HOME" value="./logs" />
<!-- 使用自定义的初始化监听器确保日志目录存在 -->
<define name="LOGBACK_DIR_CHECK" class="ch.qos.logback.core.property.FileExistsPropertyDefiner">
<path>${LOG_HOME}</path>
<createIfMissing>true</createIfMissing>
</define>
<!-- 引入Spring Boot默认配置 -->
<include resource="org/springframework/boot/logging/logback/defaults.xml" />
@@ -21,7 +21,7 @@
WHERE mac_address = #{macAddress} AND target_mac = #{targetMac}
</update>
<insert id="insert">
<insert id="insertAddressBook">
INSERT INTO ai_device_address_book (mac_address, target_mac, alias, has_permission, creator, create_date, updater, update_date)
VALUES (#{macAddress}, #{targetMac}, #{alias}, #{hasPermission}, #{creator}, NOW(), #{updater}, NOW())
</insert>
@@ -0,0 +1,57 @@
package xiaozhi.common.redis;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertInstanceOf;
import static org.junit.jupiter.api.Assertions.assertNotNull;
import java.util.ArrayList;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.springframework.data.redis.serializer.RedisSerializer;
import xiaozhi.common.utils.JsonUtils;
import xiaozhi.modules.sys.vo.SysDictDataItem;
class RedisSerializationTest {
private final RedisSerializer<Object> serializer = RedisSerializer.json();
@Test
void serverConfigRoundTripRestoresStringKeyedMaps() {
Map<String, Object> features = new HashMap<>();
features.put("addressBook", Map.of("enabled", true));
Map<String, Object> config = new HashMap<>();
config.put("features", features);
Object restored = roundTrip(config);
Map<String, Object> restoredConfig = JsonUtils.toStringObjectMap(restored);
Map<String, Object> restoredFeatures = JsonUtils.toStringObjectMap(restoredConfig.get("features"));
Map<String, Object> addressBook = JsonUtils.toStringObjectMap(restoredFeatures.get("addressBook"));
assertEquals(true, addressBook.get("enabled"));
}
@Test
void dictionaryListRoundTripRestoresDtoElements() {
SysDictDataItem item = new SysDictDataItem();
item.setName("enabled");
item.setKey("1");
Object restored = roundTrip(new ArrayList<>(List.of(item)));
List<SysDictDataItem> restoredItems = JsonUtils.toList(restored, SysDictDataItem.class);
assertInstanceOf(SysDictDataItem.class, restoredItems.get(0));
assertEquals("enabled", restoredItems.get(0).getName());
assertEquals("1", restoredItems.get(0).getKey());
}
private Object roundTrip(Object value) {
byte[] bytes = serializer.serialize(value);
assertNotNull(bytes);
Object restored = serializer.deserialize(bytes);
assertNotNull(restored);
return restored;
}
}
@@ -22,6 +22,7 @@ import java.util.function.BiFunction;
import org.apache.ibatis.binding.MapperMethod;
import org.apache.ibatis.executor.BatchResult;
import org.apache.ibatis.logging.nologging.NoLoggingImpl;
import org.apache.ibatis.mapping.MappedStatement;
import org.apache.ibatis.session.ExecutorType;
import org.apache.ibatis.session.SqlSession;
@@ -57,7 +58,8 @@ class BaseServiceImplTest {
when(sqlSessionFactory.openSession(ExecutorType.BATCH)).thenReturn(sqlSession);
when(sqlSession.insert(eq(INSERT_STATEMENT), any(TestEntity.class))).thenReturn(1);
when(sqlSession.flushStatements())
.thenReturn(List.of(batchResult(1, 1)), List.of(batchResult(1)));
.thenReturn(List.of(batchResult(1, 1)))
.thenReturn(List.of(batchResult(1)));
TestEntity first = new TestEntity(1L);
TestEntity second = new TestEntity(2L);
@@ -126,6 +128,19 @@ class BaseServiceImplTest {
verifyNoInteractions(sqlSessionFactory);
}
@Test
void deleteBatchIdsDelegatesToCompatibleApiAndPreservesAffectedRowResult() {
TestMapper mapper = mock(TestMapper.class);
List<Long> ids = List.of(1L, 2L);
service.baseDao = mapper;
when(mapper.deleteByIds(ids)).thenReturn(2).thenReturn(0);
assertTrue(service.deleteBatchIds(ids));
assertFalse(service.deleteBatchIds(ids));
verify(mapper, times(2)).deleteByIds(ids);
}
@Test
void activeTransactionSynchronizationUsesTransactionAwareCommitAndLifecycle() {
SqlSessionFactory sqlSessionFactory = mock(SqlSessionFactory.class);
@@ -170,6 +185,10 @@ class BaseServiceImplTest {
}
private static class TestService extends BaseServiceImpl<TestMapper, TestEntity> {
private TestService() {
log = new NoLoggingImpl(TestService.class.getName());
}
@Override
protected String getSqlStatement(SqlMethod sqlMethod) {
return switch (sqlMethod) {
@@ -0,0 +1,77 @@
package xiaozhi.common.utils;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertInstanceOf;
import static org.junit.jupiter.api.Assertions.assertNull;
import static org.junit.jupiter.api.Assertions.assertSame;
import static org.junit.jupiter.api.Assertions.assertThrows;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
class JsonUtilsTest {
@Test
void parsesTypedMapsAndPreservesNestedCollectionShapes() {
Map<String, Object> map = JsonUtils.parseMap(
"{\"enabled\":true,\"nested\":{\"items\":[{\"name\":\"tool\"}]}}");
Map<?, ?> nested = assertInstanceOf(Map.class, map.get("nested"));
List<?> items = assertInstanceOf(List.class, nested.get("items"));
Map<?, ?> item = assertInstanceOf(Map.class, items.get(0));
assertEquals(true, map.get("enabled"));
assertEquals("tool", item.get("name"));
List<Map<String, Object>> maps = JsonUtils.parseMapList(
"[{\"name\":\"first\",\"values\":[1,2]},{\"name\":\"second\"}]");
assertEquals("first", maps.get(0).get("name"));
assertEquals(List.of(1, 2), maps.get(0).get("values"));
assertEquals("second", maps.get(1).get("name"));
}
@Test
void checkedConvertersAreShallowMutableCopies() {
List<Object> nested = new ArrayList<>(List.of(Map.of("name", "tool")));
Map<String, Object> source = new LinkedHashMap<>();
source.put("nested", nested);
Map<String, Object> map = JsonUtils.toStringObjectMap(source);
List<Map<String, Object>> maps = JsonUtils.toStringObjectMapList(List.of(source));
List<String> strings = JsonUtils.toList(Arrays.asList("a", null), String.class);
assertSame(nested, map.get("nested"));
assertSame(nested, maps.get(0).get("nested"));
map.put("enabled", true);
maps.add(Map.of("name", "second"));
strings.add("b");
assertEquals(true, map.get("enabled"));
assertEquals(2, maps.size());
assertEquals(Arrays.asList("a", null, "b"), strings);
}
@Test
void preservesNullInputs() {
assertNull(JsonUtils.parseMap(""));
assertNull(JsonUtils.parseMapList(null));
assertNull(JsonUtils.toStringObjectMap(null));
assertNull(JsonUtils.toStringObjectMapList(null));
assertNull(JsonUtils.toList(null, String.class));
}
@Test
void rejectsUnexpectedJsonShapesAndRuntimeTypes() {
assertThrows(RuntimeException.class, () -> JsonUtils.parseMap("[]"));
assertThrows(RuntimeException.class, () -> JsonUtils.parseMapList("{}"));
assertThrows(ClassCastException.class, () -> JsonUtils.toStringObjectMap(List.of()));
assertThrows(ClassCastException.class, () -> JsonUtils.toStringObjectMap(Map.of(1, "value")));
assertThrows(ClassCastException.class, () -> JsonUtils.toStringObjectMapList(List.of("value")));
assertThrows(ClassCastException.class, () -> JsonUtils.toStringObjectMapList(Map.of()));
assertThrows(ClassCastException.class, () -> JsonUtils.toList(List.of(1), String.class));
assertThrows(ClassCastException.class, () -> JsonUtils.toList(Map.of(), String.class));
}
}
@@ -0,0 +1,96 @@
package xiaozhi.modules.config.service.impl;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertInstanceOf;
import static org.junit.jupiter.api.Assertions.assertNotSame;
import static org.junit.jupiter.api.Assertions.assertSame;
import static org.mockito.ArgumentMatchers.anyMap;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.when;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.springframework.test.util.ReflectionTestUtils;
import xiaozhi.common.redis.RedisKeys;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.modules.agent.dao.AgentVoicePrintDao;
import xiaozhi.modules.agent.service.AgentContextProviderService;
import xiaozhi.modules.agent.service.AgentMcpAccessPointService;
import xiaozhi.modules.agent.service.AgentPluginMappingService;
import xiaozhi.modules.agent.service.AgentService;
import xiaozhi.modules.agent.service.AgentTemplateService;
import xiaozhi.modules.correctword.service.CorrectWordFileService;
import xiaozhi.modules.device.service.DeviceService;
import xiaozhi.modules.model.service.ModelConfigService;
import xiaozhi.modules.sys.dto.SysParamsDTO;
import xiaozhi.modules.sys.service.SysParamsService;
import xiaozhi.modules.timbre.service.TimbreService;
import xiaozhi.modules.voiceclone.service.VoiceCloneService;
class ConfigServiceImplTest {
@Test
void cachedServerConfigIsCheckedAsAStringKeyedMapWithoutChangingNestedValues() {
RedisUtils redisUtils = mock(RedisUtils.class);
Map<String, Object> nested = new HashMap<>();
nested.put("enabled", true);
Map<String, Object> cached = new HashMap<>();
cached.put("features", nested);
when(redisUtils.get(RedisKeys.getServerConfigKey())).thenReturn(cached);
ConfigServiceImpl service = newService(mock(SysParamsService.class), redisUtils);
Map<String, Object> result = service.getConfig(true);
assertNotSame(cached, result);
assertSame(nested, result.get("features"));
assertEquals(true, ((Map<?, ?>) result.get("features")).get("enabled"));
}
@Test
void nestedSystemParametersStillShareAndPopulateTheSameConfigBranch() {
SysParamsService sysParamsService = mock(SysParamsService.class);
SysParamsDTO enabled = parameter("server.features.enabled", "true", "boolean");
SysParamsDTO labels = parameter("server.features.labels", "first;second", "array");
when(sysParamsService.list(anyMap())).thenReturn(List.of(enabled, labels));
ConfigServiceImpl service = newService(sysParamsService, mock(RedisUtils.class));
Map<String, Object> config = new HashMap<>();
Object returned = ReflectionTestUtils.invokeMethod(service, "buildConfig", config);
assertSame(config, returned);
Map<?, ?> server = assertInstanceOf(Map.class, config.get("server"));
Map<?, ?> features = assertInstanceOf(Map.class, server.get("features"));
assertEquals(true, features.get("enabled"));
assertEquals(List.of("first", "second"), features.get("labels"));
}
private static SysParamsDTO parameter(String code, String value, String type) {
SysParamsDTO parameter = new SysParamsDTO();
parameter.setParamCode(code);
parameter.setParamValue(value);
parameter.setValueType(type);
return parameter;
}
private static ConfigServiceImpl newService(SysParamsService sysParamsService, RedisUtils redisUtils) {
return new ConfigServiceImpl(
sysParamsService,
mock(DeviceService.class),
mock(ModelConfigService.class),
mock(AgentService.class),
mock(AgentTemplateService.class),
redisUtils,
mock(TimbreService.class),
mock(AgentPluginMappingService.class),
mock(AgentMcpAccessPointService.class),
mock(AgentContextProviderService.class),
mock(VoiceCloneService.class),
mock(AgentVoicePrintDao.class),
mock(CorrectWordFileService.class));
}
}
@@ -1,7 +1,6 @@
package xiaozhi.modules.device;
import java.util.HashMap;
import java.util.UUID;
import org.junit.jupiter.api.Assertions;
import org.junit.jupiter.api.DisplayName;
@@ -11,6 +10,8 @@ import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.test.context.ActiveProfiles;
import lombok.extern.slf4j.Slf4j;
import xiaozhi.common.exception.ErrorCode;
import xiaozhi.common.exception.RenException;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.modules.sys.dto.SysUserDTO;
import xiaozhi.modules.sys.service.SysUserService;
@@ -27,11 +28,13 @@ public class DeviceTest {
private SysUserService sysUserService;
@Test
public void testSaveUser() {
public void testRejectWeakPassword() {
SysUserDTO userDTO = new SysUserDTO();
userDTO.setUsername("test");
userDTO.setPassword(UUID.randomUUID().toString());
sysUserService.save(userDTO);
userDTO.setPassword("weak-password-123");
RenException exception = Assertions.assertThrows(RenException.class, () -> sysUserService.save(userDTO));
Assertions.assertEquals(ErrorCode.PASSWORD_WEAK_ERROR, exception.getCode());
}
@Test
@@ -0,0 +1,78 @@
package xiaozhi.modules.device.controller;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.mockStatic;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
import org.junit.jupiter.api.DisplayName;
import org.junit.jupiter.api.Test;
import org.mockito.MockedStatic;
import xiaozhi.common.exception.ErrorCode;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.common.user.UserDetail;
import xiaozhi.common.utils.MessageUtils;
import xiaozhi.common.utils.Result;
import xiaozhi.modules.device.dto.DeviceUpdateDTO;
import xiaozhi.modules.device.entity.DeviceEntity;
import xiaozhi.modules.device.service.DeviceAddressBookService;
import xiaozhi.modules.device.service.DeviceService;
import xiaozhi.modules.security.user.SecurityUser;
import xiaozhi.modules.sys.service.SysParamsService;
@DisplayName("设备更新接口回归测试")
class DeviceControllerTest {
private static final String DEVICE_ID = "device-id";
private static final long USER_ID = 1L;
@Test
@DisplayName("数据库未更新时不误报自动升级状态修改成功")
void updateFailureIsReturnedToCaller() {
DeviceService deviceService = mock(DeviceService.class);
DeviceEntity entity = ownedDevice();
when(deviceService.selectById(DEVICE_ID)).thenReturn(entity);
when(deviceService.updateById(entity)).thenReturn(false);
DeviceController controller = controller(deviceService);
DeviceUpdateDTO update = new DeviceUpdateDTO();
update.setAutoUpdate(0);
try (MockedStatic<SecurityUser> securityUser = mockStatic(SecurityUser.class);
MockedStatic<MessageUtils> messageUtils = mockStatic(MessageUtils.class)) {
securityUser.when(SecurityUser::getUser).thenReturn(currentUser());
messageUtils.when(() -> MessageUtils.getMessage(ErrorCode.UPDATE_DATA_FAILED))
.thenReturn("Failed to update data");
Result<Void> result = controller.updateDeviceInfo(DEVICE_ID, update);
assertEquals(ErrorCode.UPDATE_DATA_FAILED, result.getCode());
assertEquals(0, entity.getAutoUpdate());
verify(deviceService).updateById(entity);
}
}
private DeviceController controller(DeviceService deviceService) {
return new DeviceController(
deviceService,
mock(DeviceAddressBookService.class),
mock(RedisUtils.class),
mock(SysParamsService.class));
}
private DeviceEntity ownedDevice() {
DeviceEntity entity = new DeviceEntity();
entity.setId(DEVICE_ID);
entity.setUserId(USER_ID);
entity.setAutoUpdate(1);
return entity;
}
private UserDetail currentUser() {
UserDetail user = new UserDetail();
user.setId(USER_ID);
return user;
}
}
@@ -0,0 +1,46 @@
package xiaozhi.modules.device.service.impl;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.never;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
import java.util.List;
import org.junit.jupiter.api.Test;
import org.mockito.ArgumentCaptor;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.modules.device.dao.DeviceAddressBookDao;
import xiaozhi.modules.device.entity.DeviceAddressBookEntity;
import xiaozhi.modules.device.service.DeviceService;
import xiaozhi.modules.sys.service.SysParamsService;
class DeviceAddressBookServiceImplTest {
@Test
void newEntryUsesTheDedicatedAddressBookInsertMapping() {
DeviceAddressBookDao addressBookDao = mock(DeviceAddressBookDao.class);
RedisUtils redisUtils = mock(RedisUtils.class);
DeviceService deviceService = mock(DeviceService.class);
SysParamsService sysParamsService = mock(SysParamsService.class);
DeviceAddressBookServiceImpl service = new DeviceAddressBookServiceImpl(
addressBookDao, redisUtils, deviceService, sysParamsService);
when(addressBookDao.selectOne(any())).thenReturn(null);
when(addressBookDao.selectList(null)).thenReturn(List.of());
service.saveOrUpdate("00:11:22:33:44:55", "00:11:22:33:44:66", "living-room", true);
ArgumentCaptor<DeviceAddressBookEntity> entityCaptor = ArgumentCaptor.forClass(DeviceAddressBookEntity.class);
verify(addressBookDao).insertAddressBook(entityCaptor.capture());
DeviceAddressBookEntity entity = entityCaptor.getValue();
assertEquals("00:11:22:33:44:55", entity.getMacAddress());
assertEquals("00:11:22:33:44:66", entity.getTargetMac());
assertEquals("living-room", entity.getAlias());
assertTrue(entity.getHasPermission());
verify(addressBookDao, never()).insert(any(DeviceAddressBookEntity.class));
}
}
@@ -0,0 +1,104 @@
package xiaozhi.modules.device.service.impl;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertNotNull;
import static org.junit.jupiter.api.Assertions.assertNull;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.verifyNoInteractions;
import static org.mockito.Mockito.when;
import org.junit.jupiter.api.DisplayName;
import org.junit.jupiter.api.Test;
import org.springframework.aop.framework.ProxyFactory;
import xiaozhi.common.constant.Constant;
import xiaozhi.modules.device.dto.DeviceReportReqDTO;
import xiaozhi.modules.device.dto.DeviceReportRespDTO;
import xiaozhi.modules.device.entity.DeviceEntity;
import xiaozhi.modules.device.service.OtaService;
import xiaozhi.modules.sys.service.SysParamsService;
@DisplayName("设备自动升级回归测试")
class DeviceAutoUpdateTest {
private static final String MAC_ADDRESS = "00:11:22:33:44:55";
private static final String BOARD_TYPE = "test-board";
private static final String CURRENT_VERSION = "1.0.0";
@Test
@DisplayName("#3299 自动升级关闭时 OTA 响应不包含固件")
void disabledAutoUpdateOmitsFirmware() {
OtaService otaService = mock(OtaService.class);
DeviceServiceImpl service = proxiedService(deviceWithAutoUpdate(0), otaService);
DeviceReportRespDTO response = service.checkDeviceActive(
MAC_ADDRESS, MAC_ADDRESS, deviceReport());
assertNull(response.getFirmware());
verifyNoInteractions(otaService);
}
@Test
@DisplayName("自动升级开启时继续执行固件查询")
void enabledAutoUpdateChecksFirmware() {
OtaService otaService = mock(OtaService.class);
when(otaService.getLatestOta(BOARD_TYPE)).thenReturn(null);
DeviceServiceImpl service = proxiedService(deviceWithAutoUpdate(1), otaService);
DeviceReportRespDTO response = service.checkDeviceActive(
MAC_ADDRESS, MAC_ADDRESS, deviceReport());
assertNotNull(response.getFirmware());
assertEquals(CURRENT_VERSION, response.getFirmware().getVersion());
assertEquals(Constant.INVALID_FIRMWARE_URL, response.getFirmware().getUrl());
verify(otaService).getLatestOta(BOARD_TYPE);
}
private DeviceServiceImpl proxiedService(DeviceEntity device, OtaService otaService) {
SysParamsService sysParamsService = mock(SysParamsService.class);
when(sysParamsService.getValue(Constant.SERVER_WEBSOCKET, true))
.thenReturn("ws://127.0.0.1:8000/xiaozhi/v1/");
when(sysParamsService.getValue(Constant.SERVER_AUTH_ENABLED, true)).thenReturn("false");
when(sysParamsService.getValue(Constant.SERVER_MQTT_GATEWAY, true)).thenReturn(null);
DeviceServiceImpl target = new DeviceServiceImpl(
null, null, sysParamsService, null, otaService, null) {
@Override
public DeviceEntity getDeviceByMacAddress(String macAddress) {
return device;
}
@Override
public void updateDeviceConnectionInfo(String agentId, String deviceId, String appVersion) {
// No-op: connection timestamps are outside this OTA decision test.
}
};
ProxyFactory proxyFactory = new ProxyFactory(target);
proxyFactory.setProxyTargetClass(true);
proxyFactory.setExposeProxy(true);
return (DeviceServiceImpl) proxyFactory.getProxy();
}
private DeviceEntity deviceWithAutoUpdate(int autoUpdate) {
DeviceEntity device = new DeviceEntity();
device.setId(MAC_ADDRESS);
device.setMacAddress(MAC_ADDRESS);
device.setBoard(BOARD_TYPE);
device.setAutoUpdate(autoUpdate);
return device;
}
private DeviceReportReqDTO deviceReport() {
DeviceReportReqDTO.Application application = new DeviceReportReqDTO.Application();
application.setVersion(CURRENT_VERSION);
DeviceReportReqDTO.BoardInfo board = new DeviceReportReqDTO.BoardInfo();
board.setType(BOARD_TYPE);
DeviceReportReqDTO report = new DeviceReportReqDTO();
report.setApplication(application);
report.setBoard(board);
return report;
}
}
@@ -76,6 +76,25 @@ class DeviceTimeSerializationTest {
() -> assertTrue(payload.path("createDate").isNull()));
}
@ParameterizedTest(name = "自动升级状态 {0}")
@ValueSource(ints = { 0, 1 })
@DisplayName("#3299 设备列表按 autoUpdate 契约返回真实开关状态")
void serializedDeviceContainsAutoUpdateState(int autoUpdate) {
DeviceEntity entity = new DeviceEntity();
entity.setAutoUpdate(autoUpdate);
DeviceServiceImpl deviceService = serviceReturning(entity);
UserShowDeviceListVO device = deviceService.getUserDeviceList(1L, "agent-id").getFirst();
ObjectMapper objectMapper = new WebMvcConfig().jackson2HttpMessageConverter().getObjectMapper();
JsonNode payload = objectMapper.valueToTree(device);
assertAll(
() -> assertEquals(autoUpdate, device.getAutoUpdate()),
() -> assertEquals(autoUpdate, payload.path("autoUpdate").asInt()),
() -> assertTrue(payload.path("otaUpgrade").isMissingNode(),
"设备列表不应继续暴露未映射的旧字段 otaUpgrade"));
}
private DeviceServiceImpl serviceReturning(DeviceEntity entity) {
return new DeviceServiceImpl(null, null, null, null, null, null) {
@Override
@@ -0,0 +1,60 @@
package xiaozhi.modules.security.config;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertNotNull;
import static org.junit.jupiter.api.Assertions.assertTrue;
import java.lang.reflect.Method;
import java.lang.reflect.Modifier;
import org.apache.shiro.session.mgt.SessionManager;
import org.apache.shiro.spring.web.ShiroFilterFactoryBean;
import org.apache.shiro.web.mgt.WebSecurityManager;
import org.junit.jupiter.api.Test;
import org.springframework.beans.factory.config.BeanDefinition;
import org.springframework.beans.factory.config.BeanPostProcessor;
import org.springframework.context.annotation.Lazy;
import org.springframework.context.annotation.Role;
import xiaozhi.modules.security.oauth2.Oauth2Realm;
import xiaozhi.modules.sys.service.SysParamsService;
class ShiroConfigTest {
@Test
void beanPostProcessorFactoriesDoNotInstantiateShiroConfigOrBusinessDependenciesEarly() throws Exception {
Method lifecycleFactory = ShiroConfig.class.getDeclaredMethod("lifecycleBeanPostProcessor");
Method filterFactory = ShiroConfig.class.getDeclaredMethod(
"shirFilter", WebSecurityManager.class, SysParamsService.class);
assertTrue(Modifier.isStatic(lifecycleFactory.getModifiers()));
assertTrue(BeanPostProcessor.class.isAssignableFrom(ShiroFilterFactoryBean.class));
assertTrue(Modifier.isStatic(filterFactory.getModifiers()));
Lazy securityManagerLazy = filterFactory.getParameters()[0].getAnnotation(Lazy.class);
Lazy sysParamsServiceLazy = filterFactory.getParameters()[1].getAnnotation(Lazy.class);
assertNotNull(securityManagerLazy);
assertNotNull(sysParamsServiceLazy);
assertTrue(securityManagerLazy.value());
assertTrue(sysParamsServiceLazy.value());
}
@Test
void authorizationAdvisorIsInfrastructureAndDefersItsSecurityManager() throws Exception {
Method advisorFactory = ShiroConfig.class.getDeclaredMethod(
"authorizationAttributeSourceAdvisor", WebSecurityManager.class);
assertTrue(Modifier.isStatic(advisorFactory.getModifiers()));
Lazy securityManagerLazy = advisorFactory.getParameters()[0].getAnnotation(Lazy.class);
assertNotNull(securityManagerLazy);
assertTrue(securityManagerLazy.value());
assertEquals(BeanDefinition.ROLE_INFRASTRUCTURE, advisorFactory.getAnnotation(Role.class).value());
}
@Test
void securityManagerRetainsTheWebSecurityContractRequiredByShiroFilter() throws Exception {
Method securityManagerFactory = ShiroConfig.class.getDeclaredMethod(
"securityManager", Oauth2Realm.class, SessionManager.class);
assertEquals(WebSecurityManager.class, securityManagerFactory.getReturnType());
}
}
@@ -1,15 +1,23 @@
package xiaozhi.modules.sys;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertThrows;
import static org.mockito.Mockito.when;
import org.junit.jupiter.api.Test;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.test.context.ActiveProfiles;
import org.springframework.test.context.bean.override.mockito.MockitoBean;
import lombok.extern.slf4j.Slf4j;
import xiaozhi.common.exception.ErrorCode;
import xiaozhi.common.exception.RenException;
import xiaozhi.modules.security.controller.LoginController;
import xiaozhi.modules.security.dto.LoginDTO;
import xiaozhi.modules.security.dto.SmsVerificationDTO;
import xiaozhi.modules.sys.dto.RetrievePasswordDTO;
import xiaozhi.modules.sys.service.SysUserService;
@Slf4j
@SpringBootTest
@@ -19,12 +27,19 @@ class loginControllerTest {
@Autowired
LoginController loginController;
@MockitoBean
SysUserService sysUserService;
@Test
public void testRegister() {
when(sysUserService.getAllowUserRegister()).thenReturn(false);
LoginDTO loginDTO = new LoginDTO();
loginDTO.setUsername("手机号码");
loginDTO.setPassword("密码");
loginController.register(loginDTO);
RenException exception = assertThrows(RenException.class, () -> loginController.register(loginDTO));
assertEquals(ErrorCode.USER_REGISTER_DISABLED, exception.getCode());
}
@Test
@@ -0,0 +1,35 @@
package xiaozhi.modules.sys.service.impl;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertNotSame;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.when;
import java.util.ArrayList;
import java.util.List;
import org.junit.jupiter.api.Test;
import xiaozhi.common.redis.RedisKeys;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.modules.sys.dao.SysUserDao;
import xiaozhi.modules.sys.vo.SysDictDataItem;
class SysDictDataServiceImplTest {
@Test
void cachedDictionaryItemsAreCheckedAndReturnedAsDtos() {
RedisUtils redisUtils = mock(RedisUtils.class);
SysDictDataItem item = new SysDictDataItem();
item.setName("enabled");
item.setKey("1");
List<SysDictDataItem> cached = new ArrayList<>(List.of(item));
when(redisUtils.get(RedisKeys.getDictDataByTypeKey("status"))).thenReturn(cached);
SysDictDataServiceImpl service = new SysDictDataServiceImpl(mock(SysUserDao.class), redisUtils);
List<SysDictDataItem> result = service.getDictDataByType("status");
assertNotSame(cached, result);
assertEquals(List.of(item), result);
}
}
@@ -0,0 +1,36 @@
package xiaozhi.modules.sys.service.impl;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
import org.junit.jupiter.api.Test;
import org.springframework.test.util.ReflectionTestUtils;
import xiaozhi.common.constant.Constant;
import xiaozhi.modules.agent.service.AgentPluginMappingService;
import xiaozhi.modules.sys.dao.SysParamsDao;
import xiaozhi.modules.sys.redis.SysParamsRedis;
class SysParamsServiceImplTest {
@Test
void disablingAddressBookStillDeletesItsSystemPlugin() {
SysParamsRedis sysParamsRedis = mock(SysParamsRedis.class);
AgentPluginMappingService pluginMappingService = mock(AgentPluginMappingService.class);
SysParamsDao sysParamsDao = mock(SysParamsDao.class);
SysParamsServiceImpl service = new SysParamsServiceImpl(sysParamsRedis, pluginMappingService);
ReflectionTestUtils.setField(service, "baseDao", sysParamsDao);
String currentConfig = "{\"features\":{\"addressBook\":{\"enabled\":true}}}";
String newConfig = "{\"features\":{\"addressBook\":{\"enabled\":false}}}";
when(sysParamsDao.getValueByCode(Constant.SYSTEM_WEB_MENU)).thenReturn(currentConfig);
when(sysParamsDao.updateValueByCode(Constant.SYSTEM_WEB_MENU, newConfig)).thenReturn(1);
service.updateSystemWebMenu(newConfig);
verify(pluginMappingService).deleteByPluginId("SYSTEM_PLUGIN_CALL_DEVICE");
verify(sysParamsDao).updateValueByCode(Constant.SYSTEM_WEB_MENU, newConfig);
verify(sysParamsRedis).set(Constant.SYSTEM_WEB_MENU, newConfig);
}
}
@@ -2,16 +2,20 @@ package xiaozhi.modules.timbre.service.impl;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertNull;
import static org.mockito.ArgumentMatchers.argThat;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.never;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
import org.junit.jupiter.api.Test;
import org.springframework.test.util.ReflectionTestUtils;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.modules.timbre.dao.TimbreDao;
import xiaozhi.modules.timbre.dto.TimbreDataDTO;
import xiaozhi.modules.timbre.entity.TimbreEntity;
import xiaozhi.modules.timbre.vo.TimbreDetailsVO;
import xiaozhi.modules.voiceclone.dao.VoiceCloneDao;
import xiaozhi.modules.voiceclone.entity.VoiceCloneEntity;
@@ -54,4 +58,74 @@ class TimbreServiceImplTest {
assertNull(service.getDefaultLanguageById("voice-id"));
}
@Test
void updateLeavesSortOutOfTheUpdateWhenRequestOmitsIt() {
TimbreDao timbreDao = mock(TimbreDao.class);
RedisUtils redisUtils = mock(RedisUtils.class);
TimbreServiceImpl service = new TimbreServiceImpl(timbreDao, mock(VoiceCloneDao.class), redisUtils);
ReflectionTestUtils.setField(service, "baseDao", timbreDao);
TimbreDataDTO dto = validTimbreData();
service.update("voice-id", dto);
verify(timbreDao, never()).selectById("voice-id");
verify(timbreDao).updateById(argThat((TimbreEntity entity) ->
"voice-id".equals(entity.getId()) && entity.getSort() == null));
verify(redisUtils).delete("timbre:details:voice-id");
}
@Test
void updateUsesExplicitSortWithoutLoadingExistingTimbre() {
TimbreDao timbreDao = mock(TimbreDao.class);
TimbreServiceImpl service = new TimbreServiceImpl(
timbreDao, mock(VoiceCloneDao.class), mock(RedisUtils.class));
ReflectionTestUtils.setField(service, "baseDao", timbreDao);
TimbreDataDTO dto = validTimbreData();
dto.setSort(0L);
service.update("voice-id", dto);
verify(timbreDao, never()).selectById("voice-id");
verify(timbreDao).updateById(argThat((TimbreEntity entity) -> entity.getSort() == 0L));
}
@Test
void saveDefaultsOmittedSortToZero() {
TimbreDao timbreDao = mock(TimbreDao.class);
TimbreServiceImpl service = new TimbreServiceImpl(
timbreDao, mock(VoiceCloneDao.class), mock(RedisUtils.class));
ReflectionTestUtils.setField(service, "baseDao", timbreDao);
service.save(validTimbreData());
verify(timbreDao).insert(argThat((TimbreEntity entity) ->
"测试音色".equals(entity.getName()) && entity.getSort() == 0L));
}
@Test
void getSupportsLegacyRowsWithNullSort() {
TimbreDao timbreDao = mock(TimbreDao.class);
RedisUtils redisUtils = mock(RedisUtils.class);
TimbreServiceImpl service = new TimbreServiceImpl(
timbreDao, mock(VoiceCloneDao.class), redisUtils);
ReflectionTestUtils.setField(service, "baseDao", timbreDao);
TimbreEntity entity = new TimbreEntity();
entity.setId("voice-id");
entity.setSort(null);
when(timbreDao.selectById("voice-id")).thenReturn(entity);
TimbreDetailsVO details = service.get("voice-id");
assertNull(details.getSort());
}
private TimbreDataDTO validTimbreData() {
TimbreDataDTO dto = new TimbreDataDTO();
dto.setLanguages("中文");
dto.setName("测试音色");
dto.setTtsModelId("TTS_Test");
dto.setTtsVoice("test-voice");
return dto;
}
}
@@ -1,4 +1,7 @@
spring:
messages:
encoding: UTF-8
basename: i18n/messages
data:
redis:
host: localhost
+1 -1
View File
@@ -69,7 +69,7 @@
"build:quickapp-webview-huawei": "uni build -p quickapp-webview-huawei",
"build:quickapp-webview-union": "uni build -p quickapp-webview-union",
"type-check": "vue-tsc --noEmit",
"test:snapshot": "node --test src/pages/agent/components/agentSnapshotUtils.test.mjs src/pages/agent/components/agentSnapshotContracts.test.mjs",
"test:snapshot": "node --test src/pages/agent/components/agentSnapshotUtils.test.mjs src/pages/agent/components/agentSnapshotContracts.test.mjs src/pages/agent/components/voicePreviewUtils.test.mjs src/pages/device/deviceTimeUtils.test.mjs",
"openapi-ts-request": "openapi-ts",
"prepare": "git init && husky",
"lint": "eslint",
+16 -2
View File
@@ -8,6 +8,7 @@ import type {
ModelOption,
PageData,
RoleTemplate,
TtsVoice,
} from './types'
import { http } from '@/http/request/alova'
@@ -89,7 +90,7 @@ export function deleteAgent(id: string) {
// 获取TTS音色列表
export function getTTSVoices(ttsModelId: string, voiceName: string = '') {
return http.Get<{ id: string, name: string }[]>(`/models/${ttsModelId}/voices`, {
return http.Get<TtsVoice[]>(`/models/${ttsModelId}/voices`, {
params: {
voiceName,
},
@@ -214,7 +215,7 @@ export function updateAgentTags(agentId: string, data) {
// 获取所有语言
export function getAllLanguage(modelId: string) {
return http.Get<{ id: string, name: string, languages: string }[]>(`/models/${modelId}/voices`, {
return http.Get<TtsVoice[]>(`/models/${modelId}/voices`, {
meta: {
ignoreAuth: false,
toast: false,
@@ -225,6 +226,19 @@ export function getAllLanguage(modelId: string) {
})
}
/**
* ID
* @param cloneId ID
*/
export function getVoiceCloneAudioId(cloneId: string) {
return http.Post<string>(`/voiceClone/audio/${cloneId}`, {}, {
meta: {
ignoreAuth: false,
toast: false,
},
})
}
// 获取智能体历史版本列表
export function getAgentSnapshots(agentId: string, params: AgentSnapshotPageParams) {
return http.Get<PageData<AgentSnapshot>>(`/agent/${agentId}/snapshots`, {
@@ -114,6 +114,14 @@ export interface CorrectWordFile {
wordCount?: number
}
export interface TtsVoice {
id: string
name: string
voiceDemo?: string | null
languages?: string | null
isClone?: boolean | null
}
// 角色模板数据类型
export interface RoleTemplate {
id: string
+1 -1
View File
@@ -7,7 +7,7 @@ export interface Device {
id: string
userId: string
macAddress: string
lastConnectedAt: string
lastConnectedAtTimestamp: string | null
autoUpdate: number
board: string
alias?: string
@@ -0,0 +1,42 @@
/** @param {Record<string, any>} voice */
export function hasVoicePreview(voice) {
return Boolean(voice?.isClone || voice?.voiceDemo || voice?.voice_demo)
}
export function createVoicePreviewRequestGate() {
let sequence = 0
return {
begin() {
sequence += 1
return sequence
},
invalidate() {
sequence += 1
},
isCurrent(requestId) {
return requestId === sequence
},
}
}
/**
* @param {{ id: string, isClone?: boolean, voiceDemo?: string | null }} voice
* @param {(cloneId: string) => Promise<string>} getCloneAudioId
* @param {string} baseUrl
*/
export async function resolveVoicePreviewUrl(voice, getCloneAudioId, baseUrl) {
if (!voice?.isClone) {
return typeof voice?.voiceDemo === 'string' ? voice.voiceDemo : ''
}
if (!voice.id) {
return ''
}
const uuid = await getCloneAudioId(voice.id)
if (!uuid) {
return ''
}
return `${baseUrl.replace(/\/+$/, '')}/voiceClone/play/${encodeURIComponent(uuid)}`
}
@@ -0,0 +1,50 @@
/* eslint-disable test/no-import-node-test -- this zero-dependency gate intentionally uses Node's built-in runner */
import assert from 'node:assert/strict'
import test from 'node:test'
import { createVoicePreviewRequestGate, hasVoicePreview, resolveVoicePreviewUrl } from './voicePreviewUtils.mjs'
test('keeps normal voice previews on their direct URL', async () => {
let cloneRequestCount = 0
const url = await resolveVoicePreviewUrl({
id: 'normal-voice',
isClone: false,
voiceDemo: 'https://cdn.example.test/normal.wav',
}, async () => {
cloneRequestCount += 1
return 'unused'
}, 'https://api.example.test')
assert.equal(url, 'https://cdn.example.test/normal.wav')
assert.equal(cloneRequestCount, 0)
})
test('uses the clone record id to obtain and construct a temporary play URL', async () => {
let requestedCloneId = ''
const url = await resolveVoicePreviewUrl({
id: 'clone-record-id',
isClone: true,
voiceDemo: 'provider-speaker-id-must-not-be-played',
}, async (cloneId) => {
requestedCloneId = cloneId
return 'temporary uuid'
}, 'https://api.example.test/')
assert.equal(requestedCloneId, 'clone-record-id')
assert.equal(url, 'https://api.example.test/voiceClone/play/temporary%20uuid')
})
test('shows a preview control for cloned voices even without voiceDemo', () => {
assert.equal(hasVoicePreview({ id: 'clone-record-id', isClone: true }), true)
assert.equal(hasVoicePreview({ id: 'normal-voice', isClone: false, voiceDemo: '' }), false)
})
test('invalidates an older request when the same voice is cancelled and retried', () => {
const gate = createVoicePreviewRequestGate()
const firstRequest = gate.begin()
gate.invalidate()
const retryRequest = gate.begin()
assert.equal(gate.isCurrent(firstRequest), false)
assert.equal(gate.isCurrent(retryRequest), true)
})
+83 -29
View File
@@ -1,12 +1,14 @@
<script lang="ts" setup>
import type { AgentDetail, ModelOption, PluginDefinition, RoleTemplate } from '@/api/agent/types'
import type { AgentDetail, ModelOption, PluginDefinition, RoleTemplate, TtsVoice } from '@/api/agent/types'
import { computed, nextTick, onMounted, ref, watch } from 'vue'
import { getAgentDetail, getAgentTags, getAllLanguage, getModelOptions, getPluginFunctions, getRoleTemplates, updateAgent } from '@/api/agent/agent'
import { getAgentDetail, getAgentTags, getAllLanguage, getModelOptions, getPluginFunctions, getRoleTemplates, getVoiceCloneAudioId, updateAgent } from '@/api/agent/agent'
import { t } from '@/i18n'
import { usePluginStore, useProvider, useSpeedPitch } from '@/store'
import { getEnvBaseUrl } from '@/utils'
import { toast } from '@/utils/toast'
import AgentSnapshotPanel from './components/AgentSnapshotPanel.vue'
import { filterTtsVoicesByLanguage, hasUsableTtsVoiceMetadata } from './components/agentSnapshotUtils.mjs'
import { createVoicePreviewRequestGate, hasVoicePreview, resolveVoicePreviewUrl } from './components/voicePreviewUtils.mjs'
defineOptions({
name: 'AgentEdit',
@@ -84,10 +86,20 @@ const modelOptions = ref<{
TTS: [],
})
interface VoiceOption {
id?: string
value: string
name: string
voiceDemo?: string | null
voice_demo?: string | null
isClone: boolean
train_status?: number
}
//
const voiceOptions = ref<any[]>([])
const voiceOptions = ref<VoiceOption[]>([])
//
const voiceDetails = ref<Record<string, any>>({})
const voiceDetails = ref<Record<string, TtsVoice>>({})
//
const reportOptions = [
@@ -139,6 +151,7 @@ interface SnapshotRestoreContext {
//
const audioRef = ref<UniApp.InnerAudioContext | null>(null)
const playingVoiceId = ref<string>('')
const voicePreviewRequestGate = createVoicePreviewRequestGate()
// 使store
const pluginStore = usePluginStore()
@@ -513,8 +526,8 @@ interface TtsSelectionState {
languageTouched: boolean
voiceTouched: boolean
optionsModelId: string
voiceOptions: any[]
voiceDetails: Record<string, any>
voiceOptions: VoiceOption[]
voiceDetails: Record<string, TtsVoice>
languageOptions: any[]
displayNames: {
tts: string
@@ -565,7 +578,7 @@ function filterVoicesByLanguage(options: VoiceSelectionOptions = {}) {
return
}
const allVoices = Object.values(voiceDetails.value) as any[]
const allVoices = Object.values(voiceDetails.value)
//
const filteredVoices = filterTtsVoicesByLanguage(allVoices, selectedTtsLanguage.value)
@@ -624,7 +637,7 @@ async function fetchAllLanguag(ttsModelId: string, options: VoiceSelectionOption
throw new Error('No TTS voice metadata is available')
}
//
voiceDetails.value = res.reduce((acc, voice) => {
voiceDetails.value = res.reduce<Record<string, TtsVoice>>((acc, voice) => {
acc[voice.id] = voice
return acc
}, {})
@@ -863,44 +876,85 @@ function onPickerCancel(type: string) {
}
//
function playAudio(voiceDemo: string, voiceId: string, event: Event) {
async function playAudio(voice: VoiceOption, event: Event) {
event.stopPropagation() //
if (!voiceDemo) {
if (!hasVoicePreview(voice)) {
return
}
//
if (playingVoiceId.value === voiceId) {
if (playingVoiceId.value === voice.value) {
stopAudio()
return
}
//
stopAudio()
const requestId = voicePreviewRequestGate.begin()
playingVoiceId.value = voice.value
//
audioRef.value = uni.createInnerAudioContext()
audioRef.value.src = voiceDemo
playingVoiceId.value = voiceId
try {
const audioUrl = await resolveVoicePreviewUrl({
id: voice.value,
isClone: voice.isClone,
voiceDemo: voice.voiceDemo || voice.voice_demo,
}, getVoiceCloneAudioId, getEnvBaseUrl())
//
audioRef.value.onEnded(() => {
//
if (!voicePreviewRequestGate.isCurrent(requestId) || playingVoiceId.value !== voice.value) {
return
}
if (!audioUrl) {
toast.error(t('voiceprint.getAudioFailed'))
playingVoiceId.value = ''
return
}
//
const audio = uni.createInnerAudioContext()
audioRef.value = audio
audio.src = audioUrl
//
audio.onEnded(() => {
if (
voicePreviewRequestGate.isCurrent(requestId)
&& audioRef.value === audio
&& playingVoiceId.value === voice.value
) {
playingVoiceId.value = ''
}
})
//
audio.onError(() => {
if (
voicePreviewRequestGate.isCurrent(requestId)
&& audioRef.value === audio
&& playingVoiceId.value === voice.value
) {
toast.error(t('voiceprint.audioPlayFailed'))
playingVoiceId.value = ''
}
})
//
audio.play()
}
catch (error) {
if (!voicePreviewRequestGate.isCurrent(requestId) || playingVoiceId.value !== voice.value) {
return
}
console.error('获取克隆音色试听地址失败:', error)
toast.error(t('voiceprint.getAudioFailed'))
playingVoiceId.value = ''
})
//
audioRef.value.onError(() => {
toast.error('音频播放失败')
playingVoiceId.value = ''
})
//
audioRef.value.play()
}
}
//
function stopAudio() {
voicePreviewRequestGate.invalidate()
if (audioRef.value) {
audioRef.value.stop()
audioRef.value.destroy()
@@ -1592,10 +1646,10 @@ onMounted(async () => {
class="flex items-center justify-between border-b border-[#f5f5f5] p-[32rpx] transition-all active:bg-[#f5f7fb]"
@click="onPickerConfirm('voiceprint', voice.value, voice.name)"
>
<text :class="`flex-1 text-[28rpx] text-[#232338] ${(voice.voiceDemo || voice.voice_demo) ? '' : 'text-center'}`">
<text :class="`flex-1 text-[28rpx] text-[#232338] ${hasVoicePreview(voice) ? '' : 'text-center'}`">
{{ voice.name }}
</text>
<view v-if="voice.voiceDemo || voice.voice_demo" class="ml-[20rpx]" @click.stop="playAudio(voice.voiceDemo || voice.voice_demo, voice.value, $event)">
<view v-if="hasVoicePreview(voice)" class="ml-[20rpx]" @click.stop="playAudio(voice, $event)">
<wd-icon
:name="playingVoiceId === voice.value ? 'pause-circle' : 'play-circle'"
size="24px"
@@ -0,0 +1,14 @@
/** @param {unknown} timestamp */
export function parseDeviceLastConnectedAtTimestamp(timestamp) {
if (typeof timestamp !== 'string' || !timestamp.trim()) {
return null
}
const milliseconds = Number(timestamp)
if (!Number.isFinite(milliseconds)) {
return null
}
const date = new Date(milliseconds)
return Number.isNaN(date.getTime()) ? null : date
}
@@ -0,0 +1,15 @@
/* eslint-disable test/no-import-node-test -- this zero-dependency gate intentionally uses Node's built-in runner */
import assert from 'node:assert/strict'
import test from 'node:test'
import { parseDeviceLastConnectedAtTimestamp } from './deviceTimeUtils.mjs'
test('parses the backend Long timestamp serialized as a string', () => {
const timestamp = '1783689702000'
assert.equal(parseDeviceLastConnectedAtTimestamp(timestamp)?.getTime(), Number(timestamp))
})
test('rejects missing and malformed device timestamps', () => {
assert.equal(parseDeviceLastConnectedAtTimestamp(null), null)
assert.equal(parseDeviceLastConnectedAtTimestamp(''), null)
assert.equal(parseDeviceLastConnectedAtTimestamp('not-a-timestamp'), null)
})
@@ -5,6 +5,7 @@ import { useMessage } from 'wot-design-uni/components/wd-message-box'
import { bindDevice, bindDeviceManual, getBindDevices, getFirmwareTypes, unbindDevice, updateDeviceAutoUpdate } from '@/api/device'
import { t } from '@/i18n'
import { toast } from '@/utils/toast'
import { parseDeviceLastConnectedAtTimestamp } from './deviceTimeUtils.mjs'
defineOptions({
name: 'DeviceManage',
@@ -131,10 +132,10 @@ function getDeviceTypeName(boardKey: string): string {
}
//
function formatTime(timeStr: string) {
if (!timeStr)
function formatTime(timestamp: string | null) {
const date = parseDeviceLastConnectedAtTimestamp(timestamp)
if (!date)
return t('device.neverConnected')
const date = new Date(timeStr)
const now = new Date()
const diff = now.getTime() - date.getTime()
@@ -410,7 +411,7 @@ defineExpose({
{{ t('device.firmwareVersion') }}{{ device.appVersion }}
</text>
<text class="block text-[28rpx] text-[#65686f] leading-[1.4]">
{{ t('device.lastConnection') }}{{ formatTime(device.lastConnectedAt) }}
{{ t('device.lastConnection') }}{{ formatTime(device.lastConnectedAtTimestamp) }}
</text>
</view>
@@ -265,7 +265,7 @@ function showAbout() {
title: t('settings.aboutApp', { appName: import.meta.env.VITE_APP_TITLE }),
content: t('settings.aboutContent', {
appName: import.meta.env.VITE_APP_TITLE,
version: '0.9.5',
version: '0.9.6',
}),
showCancel: false,
confirmText: t('common.confirm'),
@@ -27,7 +27,7 @@ export default {
getFileList(params, callback) {
const queryParams = new URLSearchParams({
page: params.page,
pageSize: params.pageSize
limit: params.pageSize
}).toString();
RequestService.sendRequest()
@@ -79,6 +79,7 @@ export default {
remark: params.remark,
referenceAudio: params.referenceAudio,
referenceText: params.referenceText,
sort: params.sort,
ttsModelId: params.ttsModelId,
ttsVoice: params.voiceCode,
voiceDemo: params.voiceDemo || ''
@@ -141,23 +141,24 @@
<p class="section-desc">{{ $t('addressBookManagement.setPermissionDesc', { count: selectedPermissions.length }) }}</p>
</div>
<div class="section-actions">
<CustomButton size="small" @click="handleCancel">{{ $t('common.cancel') }}</CustomButton>
<CustomButton size="small" @click="handleToggleSelectAll">{{ isAllSelected ? $t('addressBookManagement.deselectAll') : $t('addressBookManagement.selectAll') }}</CustomButton>
<CustomButton type="confirm" size="small" @click="handleSavePermissions">{{ $t('addressBookManagement.save') }}</CustomButton>
<CustomButton size="small" :disabled="permissionsLoading" @click="handleCancel">{{ $t('common.cancel') }}</CustomButton>
<CustomButton size="small" :disabled="permissionsLoading" @click="handleToggleSelectAll">{{ isAllSelected ? $t('addressBookManagement.deselectAll') : $t('addressBookManagement.selectAll') }}</CustomButton>
<CustomButton type="confirm" size="small" :disabled="permissionsLoading" @click="handleSavePermissions">{{ $t('addressBookManagement.save') }}</CustomButton>
</div>
</div>
<div class="permission-grid">
<div v-loading="permissionsLoading" class="permission-grid">
<div
v-for="device in allDevices"
:key="device.id"
class="permission-item"
:class="{ active: selectedPermissions.includes(device.id) }"
:class="{ active: selectedPermissions.includes(device.deviceId) }"
>
<el-checkbox
class="permission-checkbox"
:value="selectedPermissions.includes(device.id)"
@change="(checked) => handlePermissionToggle(device.id, checked)"
:disabled="permissionsLoading"
:value="selectedPermissions.includes(device.deviceId)"
@change="(checked) => handlePermissionToggle(device.deviceId, checked)"
></el-checkbox>
<div class="permission-avatar">
<img :src="getDeviceAvatar(device.id)" alt="avatar" />
@@ -229,7 +230,9 @@ export default {
editAgentNameValue: '',
editingDeviceId: null,
editingDeviceName: '',
mqttServiceAvailable: false
mqttServiceAvailable: false,
permissionRequestSequence: 0,
permissionsLoading: false
};
},
created() {
@@ -363,18 +366,45 @@ export default {
this.loadAddressBookPermissions(device.deviceId);
},
loadAddressBookPermissions(macAddress) {
const requestId = ++this.permissionRequestSequence;
this.permissionsLoading = true;
this.selectedPermissions = [];
this.originalPermissions = [];
this.editingDeviceId = null;
this.editingDeviceName = '';
this.allDevices.forEach(device => {
device.addressBookAlias = '';
});
AddressBookApi.getAddressBookList(macAddress, (res) => {
if (
requestId !== this.permissionRequestSequence ||
this.selectedDevice?.deviceId !== macAddress
) {
return;
}
this.permissionsLoading = false;
if (res.data?.code === 0) {
const permissions = res.data.data || [];
const permissionsByTargetMac = new Map(
permissions.map(permission => [
(permission.targetMac || '').toLowerCase(),
permission
])
);
//
this.selectedPermissions = permissions
.filter(p => p.hasPermission)
.map(p => p.targetMac);
const permittedTargetMacs = new Set(
permissions
.filter(p => p.hasPermission)
.map(p => (p.targetMac || '').toLowerCase())
);
this.selectedPermissions = this.allDevices
.filter(device => permittedTargetMacs.has((device.deviceId || '').toLowerCase()))
.map(device => device.deviceId);
//
this.originalPermissions = [...this.selectedPermissions];
//
this.allDevices.forEach(device => {
const addrBook = permissions.find(p => p.targetMac === device.deviceId);
const addrBook = permissionsByTargetMac.get((device.deviceId || '').toLowerCase());
if (addrBook) {
device.addressBookAlias = addrBook.alias || '';
} else {
@@ -385,6 +415,7 @@ export default {
});
},
handleStartEditPermission(device) {
if (this.permissionsLoading) return;
this.editingDeviceId = device.id;
this.editingDeviceName = device.addressBookAlias || device.name;
this.$nextTick(() => {
@@ -412,13 +443,14 @@ export default {
this.editingDeviceId = null;
this.editingDeviceName = '';
},
handlePermissionToggle(deviceId, checked) {
handlePermissionToggle(targetMac, checked) {
if (this.permissionsLoading) return;
if (checked) {
if (!this.selectedPermissions.includes(deviceId)) {
this.selectedPermissions.push(deviceId);
if (!this.selectedPermissions.includes(targetMac)) {
this.selectedPermissions.push(targetMac);
}
} else {
const index = this.selectedPermissions.indexOf(deviceId);
const index = this.selectedPermissions.indexOf(targetMac);
if (index > -1) {
this.selectedPermissions.splice(index, 1);
}
@@ -428,21 +460,22 @@ export default {
if (this.isAllSelected) {
this.selectedPermissions = [];
} else {
this.selectedPermissions = this.allDevices.map(d => d.id);
this.selectedPermissions = this.allDevices.map(d => d.deviceId);
}
},
handleCancel() {
this.selectedPermissions = [];
},
handleSavePermissions() {
if (this.permissionsLoading) return;
const promises = this.allDevices
.filter(device => {
const isNowSelected = this.selectedPermissions.includes(device.id);
const wasOriginallySelected = this.originalPermissions.includes(device.id);
const isNowSelected = this.selectedPermissions.includes(device.deviceId);
const wasOriginallySelected = this.originalPermissions.includes(device.deviceId);
return isNowSelected !== wasOriginallySelected;
})
.map(device => {
const hasPermission = this.selectedPermissions.includes(device.id);
const hasPermission = this.selectedPermissions.includes(device.deviceId);
return new Promise((resolve) => {
AddressBookApi.updatePermission({
macAddress: this.selectedDevice.deviceId,
@@ -0,0 +1,63 @@
import assert from 'node:assert/strict';
import { readFile } from 'node:fs/promises';
import test from 'node:test';
const addressBookSource = await readFile(
new URL('../src/views/AddressBookManagement.vue', import.meta.url),
'utf8',
);
const correctWordApiSource = await readFile(
new URL('../src/apis/module/correctWord.js', import.meta.url),
'utf8',
);
test('address-book permission state consistently uses the target device MAC', () => {
assert.match(
addressBookSource,
/:value="selectedPermissions\.includes\(device\.deviceId\)"/,
);
assert.match(
addressBookSource,
/@change="\(checked\) => handlePermissionToggle\(device\.deviceId, checked\)"/,
);
assert.match(
addressBookSource,
/this\.selectedPermissions = this\.allDevices\.map\(d => d\.deviceId\)/,
);
assert.match(
addressBookSource,
/this\.originalPermissions\.includes\(device\.deviceId\)/,
);
assert.doesNotMatch(
addressBookSource,
/selectedPermissions\.includes\(device\.id\)/,
);
assert.doesNotMatch(
addressBookSource,
/originalPermissions\.includes\(device\.id\)/,
);
assert.match(
addressBookSource,
/requestId !== this\.permissionRequestSequence/,
);
assert.match(
addressBookSource,
/this\.selectedDevice\?\.deviceId !== macAddress/,
);
assert.match(
addressBookSource,
/this\.permissionsLoading = true;\s*this\.selectedPermissions = \[\];\s*this\.originalPermissions = \[\];/,
);
assert.match(
addressBookSource,
/handleSavePermissions\(\) \{\s*if \(this\.permissionsLoading\) return;/,
);
});
test('correct-word pagination maps the UI page size to the backend limit query', () => {
assert.match(
correctWordApiSource,
/new URLSearchParams\(\{\s*page: params\.page,\s*limit: params\.pageSize\s*\}\)/,
);
assert.doesNotMatch(correctWordApiSource, /pageSize: params\.pageSize/);
});
@@ -0,0 +1,17 @@
import assert from 'node:assert/strict';
import { readFile } from 'node:fs/promises';
import test from 'node:test';
const timbreApiSource = await readFile(
new URL('../src/apis/module/timbre.js', import.meta.url),
'utf8',
);
test('timbre update sends the current sort value to the backend', () => {
const updateStart = timbreApiSource.indexOf('updateVoice(params, callback)');
assert.notEqual(updateStart, -1);
const updateSource = timbreApiSource.slice(updateStart);
assert.match(updateSource, /\.method\('PUT'\)/);
assert.match(updateSource, /sort:\s*params\.sort/);
});
+1 -1
View File
@@ -6,7 +6,7 @@ from config.config_loader import load_config
from config.settings import check_config_file
from core.utils.cache.manager import cache_manager, CacheType
SERVER_VERSION = "0.9.5"
SERVER_VERSION = "0.9.6"
_logger_initialized = False
@@ -118,8 +118,17 @@ async def process_intent_result(
请根据以上信息回答用户的问题{original_text}"""
response = conn.intent.replyResult(context_prompt, original_text)
speak_txt(conn, response)
# 使用异步调用避免阻塞事件循环,影响其他设备的音频播放
try:
response = asyncio.run_coroutine_threadsafe(
conn.intent.replyResult(context_prompt, original_text),
conn.loop,
).result()
except Exception as e:
conn.logger.bind(tag=TAG).error(f"LLM生成回复失败: {e}")
response = None
if response:
speak_txt(conn, response)
conn.executor.submit(process_context_result)
return True
@@ -184,7 +193,15 @@ async def process_intent_result(
elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复
text = result.result
conn.dialogue.put(Message(role="tool", content=text))
llm_result = conn.intent.replyResult(text, original_text)
# 使用异步调用避免阻塞事件循环,影响其他设备的音频播放
try:
llm_result = asyncio.run_coroutine_threadsafe(
conn.intent.replyResult(text, original_text),
conn.loop,
).result()
except Exception as e:
conn.logger.bind(tag=TAG).error(f"LLM生成回复失败: {e}")
llm_result = text
if llm_result is None:
llm_result = text
speak_txt(conn, llm_result)
@@ -1,3 +1,4 @@
import asyncio
from typing import List, Dict, TYPE_CHECKING
if TYPE_CHECKING:
@@ -120,12 +121,18 @@ class IntentProvider(IntentProviderBase):
)
return prompt
def replyResult(self, text: str, original_text: str):
async def replyResult(self, text: str, original_text: str):
"""使用 asyncio.to_thread 避免阻塞事件循环"""
try:
llm_result = self.llm.response_no_stream(
user_prompt = (
"请根据以上内容,像人类一样说话的口吻回复用户,要求简洁,请直接返回结果。用户现在说:"
+ original_text
)
# 使用 to_thread 将同步阻塞调用放到线程池中执行,不阻塞事件循环
llm_result = await asyncio.to_thread(
self.llm.response_no_stream,
system_prompt=text,
user_prompt="请根据以上内容,像人类一样说话的口吻回复用户,要求简洁,请直接返回结果。用户现在说:"
+ original_text,
user_prompt=user_prompt,
)
return llm_result
except Exception as e:
@@ -207,8 +214,11 @@ class IntentProvider(IntentProviderBase):
logger.bind(tag=TAG).debug(f"开始LLM意图识别调用, 模型: {model_info}")
try:
intent = self.llm.response_no_stream(
system_prompt=prompt_music, user_prompt=user_prompt
# 使用 to_thread 将同步阻塞调用放到线程池中,避免阻塞事件循环
intent = await asyncio.to_thread(
self.llm.response_no_stream,
system_prompt=prompt_music,
user_prompt=user_prompt,
)
except Exception as e:
logger.bind(tag=TAG).error(f"Error in intent detection LLM call: {e}")
@@ -1,6 +1,7 @@
import random
import httpx
from markitdown import MarkItDown
from io import BytesIO
from markitdown import MarkItDown, StreamInfo
from config.logger import setup_logging
from plugins_func.register import register_function, ToolType, ActionResponse, Action
from typing import TYPE_CHECKING
@@ -144,7 +145,14 @@ async def fetch_news_detail(url):
# 使用MarkItDown清理HTML内容
md = MarkItDown(enable_plugins=False)
result = md.convert(response)
result = md.convert_stream(
BytesIO(response.content),
stream_info=StreamInfo(
mimetype="text/html",
extension=".html",
charset=response.encoding or "utf-8",
),
)
# 获取清理后的文本内容
clean_text = result.text_content