mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-24 08:03:53 +08:00
Merge remote-tracking branch 'origin/main' into fix-agent-idor
# Conflicts: # main/manager-api/src/main/java/xiaozhi/modules/agent/controller/AgentController.java # main/manager-api/src/main/java/xiaozhi/modules/agent/service/AgentService.java # main/manager-api/src/main/java/xiaozhi/modules/agent/service/impl/AgentServiceImpl.java
This commit is contained in:
@@ -0,0 +1,103 @@
|
||||
package xiaozhi.modules.agent.Enums;
|
||||
|
||||
import java.util.Arrays;
|
||||
import java.util.List;
|
||||
import java.util.function.BiConsumer;
|
||||
import java.util.function.Function;
|
||||
|
||||
import lombok.Getter;
|
||||
import xiaozhi.modules.agent.dto.AgentSnapshotDataDTO;
|
||||
import xiaozhi.modules.agent.dto.AgentUpdateDTO;
|
||||
import xiaozhi.modules.agent.entity.AgentEntity;
|
||||
|
||||
@Getter
|
||||
public enum AgentSnapshotField {
|
||||
AGENT_CODE("agentCode", AgentSnapshotDataDTO::getAgentCode, AgentUpdateDTO::getAgentCode,
|
||||
(agent, data) -> agent.setAgentCode(data.getAgentCode())),
|
||||
AGENT_NAME("agentName", AgentSnapshotDataDTO::getAgentName, AgentUpdateDTO::getAgentName,
|
||||
(agent, data) -> agent.setAgentName(data.getAgentName())),
|
||||
ASR_MODEL_ID("asrModelId", AgentSnapshotDataDTO::getAsrModelId, AgentUpdateDTO::getAsrModelId,
|
||||
(agent, data) -> agent.setAsrModelId(data.getAsrModelId())),
|
||||
VAD_MODEL_ID("vadModelId", AgentSnapshotDataDTO::getVadModelId, AgentUpdateDTO::getVadModelId,
|
||||
(agent, data) -> agent.setVadModelId(data.getVadModelId())),
|
||||
LLM_MODEL_ID("llmModelId", AgentSnapshotDataDTO::getLlmModelId, AgentUpdateDTO::getLlmModelId,
|
||||
(agent, data) -> agent.setLlmModelId(data.getLlmModelId())),
|
||||
SLM_MODEL_ID("slmModelId", AgentSnapshotDataDTO::getSlmModelId, AgentUpdateDTO::getSlmModelId,
|
||||
(agent, data) -> agent.setSlmModelId(data.getSlmModelId())),
|
||||
VLLM_MODEL_ID("vllmModelId", AgentSnapshotDataDTO::getVllmModelId, AgentUpdateDTO::getVllmModelId,
|
||||
(agent, data) -> agent.setVllmModelId(data.getVllmModelId())),
|
||||
TTS_MODEL_ID("ttsModelId", AgentSnapshotDataDTO::getTtsModelId, AgentUpdateDTO::getTtsModelId,
|
||||
(agent, data) -> agent.setTtsModelId(data.getTtsModelId())),
|
||||
TTS_VOICE_ID("ttsVoiceId", AgentSnapshotDataDTO::getTtsVoiceId, AgentUpdateDTO::getTtsVoiceId,
|
||||
(agent, data) -> agent.setTtsVoiceId(data.getTtsVoiceId())),
|
||||
TTS_LANGUAGE("ttsLanguage", AgentSnapshotDataDTO::getTtsLanguage, AgentUpdateDTO::getTtsLanguage,
|
||||
(agent, data) -> agent.setTtsLanguage(data.getTtsLanguage())),
|
||||
TTS_VOLUME("ttsVolume", AgentSnapshotDataDTO::getTtsVolume, AgentUpdateDTO::getTtsVolume,
|
||||
(agent, data) -> agent.setTtsVolume(data.getTtsVolume())),
|
||||
TTS_RATE("ttsRate", AgentSnapshotDataDTO::getTtsRate, AgentUpdateDTO::getTtsRate,
|
||||
(agent, data) -> agent.setTtsRate(data.getTtsRate())),
|
||||
TTS_PITCH("ttsPitch", AgentSnapshotDataDTO::getTtsPitch, AgentUpdateDTO::getTtsPitch,
|
||||
(agent, data) -> agent.setTtsPitch(data.getTtsPitch())),
|
||||
MEM_MODEL_ID("memModelId", AgentSnapshotDataDTO::getMemModelId, AgentUpdateDTO::getMemModelId,
|
||||
(agent, data) -> agent.setMemModelId(data.getMemModelId())),
|
||||
INTENT_MODEL_ID("intentModelId", AgentSnapshotDataDTO::getIntentModelId, AgentUpdateDTO::getIntentModelId,
|
||||
(agent, data) -> agent.setIntentModelId(data.getIntentModelId())),
|
||||
CHAT_HISTORY_CONF("chatHistoryConf", AgentSnapshotDataDTO::getChatHistoryConf, AgentUpdateDTO::getChatHistoryConf,
|
||||
(agent, data) -> agent.setChatHistoryConf(data.getChatHistoryConf())),
|
||||
SYSTEM_PROMPT("systemPrompt", AgentSnapshotDataDTO::getSystemPrompt, AgentUpdateDTO::getSystemPrompt,
|
||||
(agent, data) -> agent.setSystemPrompt(data.getSystemPrompt())),
|
||||
SUMMARY_MEMORY("summaryMemory", AgentSnapshotDataDTO::getSummaryMemory, AgentUpdateDTO::getSummaryMemory,
|
||||
(agent, data) -> agent.setSummaryMemory(data.getSummaryMemory())),
|
||||
LANG_CODE("langCode", AgentSnapshotDataDTO::getLangCode, AgentUpdateDTO::getLangCode,
|
||||
(agent, data) -> agent.setLangCode(data.getLangCode())),
|
||||
LANGUAGE("language", AgentSnapshotDataDTO::getLanguage, AgentUpdateDTO::getLanguage,
|
||||
(agent, data) -> agent.setLanguage(data.getLanguage())),
|
||||
SORT("sort", AgentSnapshotDataDTO::getSort, AgentUpdateDTO::getSort,
|
||||
(agent, data) -> agent.setSort(data.getSort())),
|
||||
FUNCTIONS("functions", AgentSnapshotDataDTO::getFunctions, AgentUpdateDTO::getFunctions, null),
|
||||
CONTEXT_PROVIDERS("contextProviders", AgentSnapshotDataDTO::getContextProviders,
|
||||
AgentUpdateDTO::getContextProviders, null),
|
||||
CORRECT_WORD_FILE_IDS("correctWordFileIds", AgentSnapshotDataDTO::getCorrectWordFileIds,
|
||||
AgentUpdateDTO::getCorrectWordFileIds, null),
|
||||
TAG_NAMES("tagNames", AgentSnapshotDataDTO::getTagNames, AgentUpdateDTO::getTagNames, null);
|
||||
|
||||
private final String fieldName;
|
||||
private final Function<AgentSnapshotDataDTO, Object> snapshotGetter;
|
||||
private final Function<AgentUpdateDTO, Object> updateGetter;
|
||||
private final BiConsumer<AgentEntity, AgentSnapshotDataDTO> restoreApplier;
|
||||
|
||||
AgentSnapshotField(String fieldName, Function<AgentSnapshotDataDTO, Object> snapshotGetter,
|
||||
Function<AgentUpdateDTO, Object> updateGetter,
|
||||
BiConsumer<AgentEntity, AgentSnapshotDataDTO> restoreApplier) {
|
||||
this.fieldName = fieldName;
|
||||
this.snapshotGetter = snapshotGetter;
|
||||
this.updateGetter = updateGetter;
|
||||
this.restoreApplier = restoreApplier;
|
||||
}
|
||||
|
||||
public static List<String> names() {
|
||||
return Arrays.stream(values()).map(AgentSnapshotField::getFieldName).toList();
|
||||
}
|
||||
|
||||
public static String canonicalName(String fieldName) {
|
||||
return "tags".equals(fieldName) ? TAG_NAMES.getFieldName() : fieldName;
|
||||
}
|
||||
|
||||
public Object snapshotValue(AgentSnapshotDataDTO data) {
|
||||
return data == null ? null : snapshotGetter.apply(data);
|
||||
}
|
||||
|
||||
public Object updateValue(AgentUpdateDTO data) {
|
||||
return data == null ? null : updateGetter.apply(data);
|
||||
}
|
||||
|
||||
public boolean isRestorableAgentField() {
|
||||
return restoreApplier != null;
|
||||
}
|
||||
|
||||
public void applyTo(AgentEntity agent, AgentSnapshotDataDTO data) {
|
||||
if (restoreApplier != null) {
|
||||
restoreApplier.accept(agent, data);
|
||||
}
|
||||
}
|
||||
}
|
||||
+4
-1
@@ -342,7 +342,10 @@ public class AgentController {
|
||||
requireAgentPermission(id);
|
||||
List<String> tagIds = (List<String>) params.get("tagIds");
|
||||
List<String> tagNames = (List<String>) params.get("tagNames");
|
||||
agentTagService.saveAgentTags(id, tagIds, tagNames);
|
||||
AgentUpdateDTO dto = new AgentUpdateDTO();
|
||||
dto.setTagIds(tagIds);
|
||||
dto.setTagNames(tagNames);
|
||||
agentService.updateAgentById(id, dto);
|
||||
return new Result<Void>().ok(null);
|
||||
}
|
||||
|
||||
|
||||
+75
@@ -0,0 +1,75 @@
|
||||
package xiaozhi.modules.agent.controller;
|
||||
|
||||
import org.apache.shiro.authz.annotation.RequiresPermissions;
|
||||
import org.springdoc.core.annotations.ParameterObject;
|
||||
import org.springframework.web.bind.annotation.DeleteMapping;
|
||||
import org.springframework.web.bind.annotation.GetMapping;
|
||||
import org.springframework.web.bind.annotation.PathVariable;
|
||||
import org.springframework.web.bind.annotation.PostMapping;
|
||||
import org.springframework.web.bind.annotation.RequestMapping;
|
||||
import org.springframework.web.bind.annotation.RestController;
|
||||
|
||||
import io.swagger.v3.oas.annotations.Operation;
|
||||
import io.swagger.v3.oas.annotations.tags.Tag;
|
||||
import lombok.AllArgsConstructor;
|
||||
import xiaozhi.common.exception.RenException;
|
||||
import xiaozhi.common.page.PageData;
|
||||
import xiaozhi.common.user.UserDetail;
|
||||
import xiaozhi.common.utils.Result;
|
||||
import xiaozhi.modules.agent.dto.AgentSnapshotPageDTO;
|
||||
import xiaozhi.modules.agent.service.AgentService;
|
||||
import xiaozhi.modules.agent.service.AgentSnapshotService;
|
||||
import xiaozhi.modules.agent.vo.AgentSnapshotVO;
|
||||
import xiaozhi.modules.security.user.SecurityUser;
|
||||
|
||||
@Tag(name = "智能体快照")
|
||||
@AllArgsConstructor
|
||||
@RestController
|
||||
@RequestMapping("/agent/{agentId}/snapshots")
|
||||
public class AgentSnapshotController {
|
||||
private final AgentSnapshotService agentSnapshotService;
|
||||
private final AgentService agentService;
|
||||
|
||||
@GetMapping
|
||||
@Operation(summary = "获取智能体快照列表")
|
||||
@RequiresPermissions("sys:role:normal")
|
||||
public Result<PageData<AgentSnapshotVO>> page(
|
||||
@PathVariable String agentId,
|
||||
@ParameterObject AgentSnapshotPageDTO params) {
|
||||
checkPermission(agentId);
|
||||
return new Result<PageData<AgentSnapshotVO>>().ok(agentSnapshotService.page(agentId, params));
|
||||
}
|
||||
|
||||
@GetMapping("/{snapshotId}")
|
||||
@Operation(summary = "获取智能体快照详情")
|
||||
@RequiresPermissions("sys:role:normal")
|
||||
public Result<AgentSnapshotVO> getSnapshot(@PathVariable String agentId, @PathVariable String snapshotId) {
|
||||
checkPermission(agentId);
|
||||
return new Result<AgentSnapshotVO>().ok(agentSnapshotService.getSnapshot(agentId, snapshotId));
|
||||
}
|
||||
|
||||
@PostMapping("/{snapshotId}/restore")
|
||||
@Operation(summary = "恢复智能体快照")
|
||||
@RequiresPermissions("sys:role:normal")
|
||||
public Result<Void> restore(@PathVariable String agentId, @PathVariable String snapshotId) {
|
||||
checkPermission(agentId);
|
||||
agentSnapshotService.restoreSnapshot(agentId, snapshotId);
|
||||
return new Result<>();
|
||||
}
|
||||
|
||||
@DeleteMapping("/{snapshotId}")
|
||||
@Operation(summary = "删除智能体历史快照")
|
||||
@RequiresPermissions("sys:role:normal")
|
||||
public Result<Void> deleteSnapshot(@PathVariable String agentId, @PathVariable String snapshotId) {
|
||||
checkPermission(agentId);
|
||||
agentSnapshotService.deleteSnapshot(agentId, snapshotId);
|
||||
return new Result<>();
|
||||
}
|
||||
|
||||
private void checkPermission(String agentId) {
|
||||
UserDetail user = SecurityUser.getUser();
|
||||
if (user == null || !agentService.checkAgentPermission(agentId, user.getId())) {
|
||||
throw new RenException("没有权限访问该智能体快照");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -36,4 +36,11 @@ public interface AgentDao extends BaseDao<AgentEntity> {
|
||||
* @param agentId 智能体ID
|
||||
*/
|
||||
AgentInfoVO selectAgentInfoById(@Param("agentId") String agentId);
|
||||
|
||||
/**
|
||||
* 锁定智能体主记录,用于串行化同一智能体的配置写入
|
||||
*
|
||||
* @param agentId 智能体ID
|
||||
*/
|
||||
AgentEntity selectByIdForUpdate(@Param("agentId") String agentId);
|
||||
}
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
package xiaozhi.modules.agent.dao;
|
||||
|
||||
import org.apache.ibatis.annotations.Mapper;
|
||||
import org.apache.ibatis.annotations.Param;
|
||||
|
||||
import xiaozhi.common.dao.BaseDao;
|
||||
import xiaozhi.modules.agent.entity.AgentSnapshotEntity;
|
||||
|
||||
@Mapper
|
||||
public interface AgentSnapshotDao extends BaseDao<AgentSnapshotEntity> {
|
||||
Integer selectMaxVersionNo(@Param("agentId") String agentId);
|
||||
|
||||
AgentSnapshotEntity selectLatestSnapshot(@Param("agentId") String agentId);
|
||||
|
||||
AgentSnapshotEntity selectNextSnapshot(@Param("agentId") String agentId, @Param("versionNo") Integer versionNo);
|
||||
|
||||
int insertWithNextVersion(@Param("snapshot") AgentSnapshotEntity snapshot);
|
||||
|
||||
int deleteOlderThanKeepLimit(@Param("agentId") String agentId, @Param("keepLimit") int keepLimit);
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
package xiaozhi.modules.agent.dto;
|
||||
|
||||
import java.io.Serializable;
|
||||
import java.util.List;
|
||||
|
||||
import io.swagger.v3.oas.annotations.media.Schema;
|
||||
import lombok.Data;
|
||||
|
||||
@Data
|
||||
@Schema(description = "智能体快照数据")
|
||||
public class AgentSnapshotDataDTO implements Serializable {
|
||||
private static final long serialVersionUID = 1L;
|
||||
|
||||
private String agentCode;
|
||||
private String agentName;
|
||||
private String asrModelId;
|
||||
private String vadModelId;
|
||||
private String llmModelId;
|
||||
private String slmModelId;
|
||||
private String vllmModelId;
|
||||
private String ttsModelId;
|
||||
private String ttsVoiceId;
|
||||
private String ttsLanguage;
|
||||
private Integer ttsVolume;
|
||||
private Integer ttsRate;
|
||||
private Integer ttsPitch;
|
||||
private String memModelId;
|
||||
private String intentModelId;
|
||||
private Integer chatHistoryConf;
|
||||
private String systemPrompt;
|
||||
private String summaryMemory;
|
||||
private String langCode;
|
||||
private String language;
|
||||
private Integer sort;
|
||||
private List<AgentUpdateDTO.FunctionInfo> functions;
|
||||
private List<ContextProviderDTO> contextProviders;
|
||||
private List<String> correctWordFileIds;
|
||||
private List<String> tagNames;
|
||||
private List<AgentSnapshotTagDTO> tags;
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
package xiaozhi.modules.agent.dto;
|
||||
|
||||
import io.swagger.v3.oas.annotations.media.Schema;
|
||||
import lombok.Data;
|
||||
|
||||
@Data
|
||||
@Schema(description = "智能体快照分页查询参数")
|
||||
public class AgentSnapshotPageDTO {
|
||||
@Schema(description = "当前页码,从1开始", example = "1")
|
||||
private Integer page = 1;
|
||||
|
||||
@Schema(description = "每页数量", example = "10")
|
||||
private Integer limit = 10;
|
||||
|
||||
@Schema(description = "版本锚点,只查询小于等于该版本号的历史快照", example = "20")
|
||||
private Integer maxVersionNo;
|
||||
|
||||
public int pageOrDefault() {
|
||||
return page == null || page < 1 ? 1 : page;
|
||||
}
|
||||
|
||||
public int limitOrDefault() {
|
||||
if (limit == null || limit < 1) {
|
||||
return 10;
|
||||
}
|
||||
return limit;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
package xiaozhi.modules.agent.dto;
|
||||
|
||||
import java.io.Serializable;
|
||||
|
||||
import io.swagger.v3.oas.annotations.media.Schema;
|
||||
import lombok.Data;
|
||||
|
||||
@Data
|
||||
@Schema(description = "智能体快照标签")
|
||||
public class AgentSnapshotTagDTO implements Serializable {
|
||||
private static final long serialVersionUID = 1L;
|
||||
|
||||
private String id;
|
||||
private String tagName;
|
||||
private Integer sort;
|
||||
}
|
||||
@@ -4,9 +4,12 @@ import java.io.Serializable;
|
||||
import java.math.BigDecimal;
|
||||
import java.util.HashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import com.fasterxml.jackson.core.type.TypeReference;
|
||||
import io.swagger.v3.oas.annotations.media.Schema;
|
||||
import lombok.Data;
|
||||
import xiaozhi.common.utils.JsonUtils;
|
||||
|
||||
/**
|
||||
* 智能体更新DTO
|
||||
@@ -91,15 +94,50 @@ public class AgentUpdateDTO implements Serializable {
|
||||
@Schema(description = "替换词文件ID列表", nullable = true)
|
||||
private List<String> correctWordFileIds;
|
||||
|
||||
@Schema(description = "标签名称列表", nullable = true)
|
||||
private List<String> tagNames;
|
||||
|
||||
@Schema(description = "标签ID列表", nullable = true)
|
||||
private List<String> tagIds;
|
||||
|
||||
@Data
|
||||
@Schema(description = "插件函数信息")
|
||||
public static class FunctionInfo implements Serializable {
|
||||
private static final TypeReference<HashMap<String, Object>> PARAM_INFO_TYPE = new TypeReference<>() {
|
||||
};
|
||||
|
||||
@Schema(description = "插件ID", example = "plugin_01")
|
||||
private String pluginId;
|
||||
|
||||
@Schema(description = "函数参数信息", nullable = true)
|
||||
private HashMap<String, Object> paramInfo;
|
||||
private HashMap<String, Object> paramInfo = new HashMap<>();
|
||||
|
||||
public void setParamInfo(Object paramInfo) {
|
||||
this.paramInfo = normalizeParamInfo(paramInfo);
|
||||
}
|
||||
|
||||
private static HashMap<String, Object> normalizeParamInfo(Object paramInfo) {
|
||||
if (paramInfo == null) {
|
||||
return new HashMap<>();
|
||||
}
|
||||
if (paramInfo instanceof String value) {
|
||||
if (value.trim().isEmpty()) {
|
||||
return new HashMap<>();
|
||||
}
|
||||
return JsonUtils.parseObject(value, PARAM_INFO_TYPE);
|
||||
}
|
||||
if (paramInfo instanceof Map<?, ?> value) {
|
||||
HashMap<String, Object> normalized = new HashMap<>();
|
||||
value.forEach((key, val) -> {
|
||||
if (key != null) {
|
||||
normalized.put(String.valueOf(key), val);
|
||||
}
|
||||
});
|
||||
return normalized;
|
||||
}
|
||||
return JsonUtils.parseObject(JsonUtils.toJsonString(paramInfo), PARAM_INFO_TYPE);
|
||||
}
|
||||
|
||||
private static final long serialVersionUID = 1L;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+2
-2
@@ -7,11 +7,11 @@ import com.baomidou.mybatisplus.annotation.IdType;
|
||||
import com.baomidou.mybatisplus.annotation.TableField;
|
||||
import com.baomidou.mybatisplus.annotation.TableId;
|
||||
import com.baomidou.mybatisplus.annotation.TableName;
|
||||
import com.baomidou.mybatisplus.extension.handlers.JacksonTypeHandler;
|
||||
|
||||
import io.swagger.v3.oas.annotations.media.Schema;
|
||||
import lombok.Data;
|
||||
import xiaozhi.modules.agent.dto.ContextProviderDTO;
|
||||
import xiaozhi.modules.agent.typehandler.ContextProviderListTypeHandler;
|
||||
|
||||
@Data
|
||||
@TableName(value = "ai_agent_context_provider", autoResultMap = true)
|
||||
@@ -26,7 +26,7 @@ public class AgentContextProviderEntity {
|
||||
private String agentId;
|
||||
|
||||
@Schema(description = "上下文源配置")
|
||||
@TableField(typeHandler = JacksonTypeHandler.class)
|
||||
@TableField(typeHandler = ContextProviderListTypeHandler.class)
|
||||
private List<ContextProviderDTO> contextProviders;
|
||||
|
||||
@Schema(description = "创建者")
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
package xiaozhi.modules.agent.entity;
|
||||
|
||||
import java.util.Date;
|
||||
|
||||
import com.baomidou.mybatisplus.annotation.IdType;
|
||||
import com.baomidou.mybatisplus.annotation.TableId;
|
||||
import com.baomidou.mybatisplus.annotation.TableName;
|
||||
|
||||
import io.swagger.v3.oas.annotations.media.Schema;
|
||||
import lombok.Data;
|
||||
|
||||
@Data
|
||||
@TableName("ai_agent_snapshot")
|
||||
@Schema(description = "智能体配置快照")
|
||||
public class AgentSnapshotEntity {
|
||||
|
||||
@TableId(type = IdType.ASSIGN_UUID)
|
||||
@Schema(description = "快照ID")
|
||||
private String id;
|
||||
|
||||
@Schema(description = "智能体ID")
|
||||
private String agentId;
|
||||
|
||||
@Schema(description = "所属用户ID")
|
||||
private Long userId;
|
||||
|
||||
@Schema(description = "版本号")
|
||||
private Integer versionNo;
|
||||
|
||||
@Schema(description = "快照数据JSON")
|
||||
private String snapshotData;
|
||||
|
||||
@Schema(description = "变更字段JSON")
|
||||
private String changedFields;
|
||||
|
||||
@Schema(description = "快照来源")
|
||||
private String source;
|
||||
|
||||
@Schema(description = "恢复来源快照ID")
|
||||
private String restoreFromSnapshotId;
|
||||
|
||||
@Schema(description = "恢复来源版本号")
|
||||
private Integer restoreFromVersionNo;
|
||||
|
||||
@Schema(description = "创建者")
|
||||
private Long creator;
|
||||
|
||||
@Schema(description = "创建时间")
|
||||
private Date createdAt;
|
||||
}
|
||||
@@ -60,6 +60,13 @@ public interface AgentService extends BaseService<AgentEntity> {
|
||||
*/
|
||||
void deleteAgentByUserId(Long userId);
|
||||
|
||||
/**
|
||||
* 删除智能体及其关联数据
|
||||
*
|
||||
* @param agentId 智能体ID
|
||||
*/
|
||||
void deleteAgent(String agentId);
|
||||
|
||||
/**
|
||||
* 获取用户智能体列表
|
||||
*
|
||||
@@ -129,6 +136,15 @@ public interface AgentService extends BaseService<AgentEntity> {
|
||||
*/
|
||||
void deleteAgentById(String agentId, Long userId);
|
||||
|
||||
/**
|
||||
* 更新智能体
|
||||
*
|
||||
* @param agentId 智能体ID
|
||||
* @param dto 更新智能体所需的信息
|
||||
* @param createSnapshot 是否创建配置快照
|
||||
*/
|
||||
void updateAgentById(String agentId, AgentUpdateDTO dto, boolean createSnapshot);
|
||||
|
||||
/**
|
||||
* 创建智能体
|
||||
*
|
||||
|
||||
+23
@@ -0,0 +1,23 @@
|
||||
package xiaozhi.modules.agent.service;
|
||||
|
||||
import xiaozhi.common.page.PageData;
|
||||
import xiaozhi.common.service.BaseService;
|
||||
import xiaozhi.modules.agent.dto.AgentSnapshotPageDTO;
|
||||
import xiaozhi.modules.agent.entity.AgentSnapshotEntity;
|
||||
import xiaozhi.modules.agent.vo.AgentSnapshotVO;
|
||||
|
||||
public interface AgentSnapshotService extends BaseService<AgentSnapshotEntity> {
|
||||
void createSnapshot(String agentId, String source);
|
||||
|
||||
PageData<AgentSnapshotVO> page(String agentId, AgentSnapshotPageDTO params);
|
||||
|
||||
AgentSnapshotVO getSnapshot(String agentId, String snapshotId);
|
||||
|
||||
void restoreSnapshot(String agentId, String snapshotId);
|
||||
|
||||
void deleteSnapshot(String agentId, String snapshotId);
|
||||
|
||||
Integer getCurrentVersionNo(String agentId);
|
||||
|
||||
void deleteByAgentId(String agentId);
|
||||
}
|
||||
+2
-2
@@ -118,7 +118,7 @@ public class AgentChatSummaryServiceImpl implements AgentChatSummaryService {
|
||||
{
|
||||
setSummaryMemory(summaryDTO.getSummary());
|
||||
}
|
||||
});
|
||||
}, false);
|
||||
log.info("成功保存会话 {} 的聊天记录总结到智能体 {}", sessionId, agentId);
|
||||
} else {
|
||||
log.info("生成总结失败: {}", summaryDTO.getErrorMessage());
|
||||
@@ -518,4 +518,4 @@ public class AgentChatSummaryServiceImpl implements AgentChatSummaryService {
|
||||
return null;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+104
-67
@@ -15,7 +15,6 @@ import org.springframework.stereotype.Service;
|
||||
import org.springframework.transaction.annotation.Transactional;
|
||||
|
||||
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
|
||||
import com.baomidou.mybatisplus.core.conditions.update.UpdateWrapper;
|
||||
import com.baomidou.mybatisplus.core.metadata.IPage;
|
||||
|
||||
import lombok.AllArgsConstructor;
|
||||
@@ -44,9 +43,10 @@ import xiaozhi.modules.agent.entity.AgentTagEntity;
|
||||
import xiaozhi.modules.agent.entity.AgentTemplateEntity;
|
||||
import xiaozhi.modules.agent.service.AgentChatHistoryService;
|
||||
import xiaozhi.modules.agent.service.AgentContextProviderService;
|
||||
import xiaozhi.modules.agent.service.AgentPluginMappingService;
|
||||
import xiaozhi.modules.agent.service.AgentService;
|
||||
import xiaozhi.modules.agent.service.AgentTagService;
|
||||
import xiaozhi.modules.agent.service.AgentPluginMappingService;
|
||||
import xiaozhi.modules.agent.service.AgentService;
|
||||
import xiaozhi.modules.agent.service.AgentSnapshotService;
|
||||
import xiaozhi.modules.agent.service.AgentTagService;
|
||||
import xiaozhi.modules.agent.service.AgentTemplateService;
|
||||
import xiaozhi.modules.agent.vo.AgentInfoVO;
|
||||
import xiaozhi.modules.correctword.service.CorrectWordFileService;
|
||||
@@ -74,9 +74,10 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
|
||||
private final AgentChatHistoryService agentChatHistoryService;
|
||||
private final AgentTemplateService agentTemplateService;
|
||||
private final ModelProviderService modelProviderService;
|
||||
private final AgentContextProviderService agentContextProviderService;
|
||||
private final AgentTagService agentTagService;
|
||||
private final CorrectWordFileService correctWordFileService;
|
||||
private final AgentContextProviderService agentContextProviderService;
|
||||
private final AgentTagService agentTagService;
|
||||
private final CorrectWordFileService correctWordFileService;
|
||||
private final AgentSnapshotService agentSnapshotService;
|
||||
|
||||
@Override
|
||||
public PageData<AgentEntity> adminAgentList(Map<String, Object> params) {
|
||||
@@ -108,12 +109,13 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
|
||||
agent.setContextProviders(contextProviderEntity.getContextProviders());
|
||||
}
|
||||
|
||||
// 查询替换词文件ID列表
|
||||
List<String> correctWordFileIds = correctWordFileService.getAgentCorrectWordFileIds(id);
|
||||
agent.setCorrectWordFileIds(correctWordFileIds);
|
||||
|
||||
// 无需额外查询插件列表,已通过SQL查询出来
|
||||
return agent;
|
||||
// 查询替换词文件ID列表
|
||||
List<String> correctWordFileIds = correctWordFileService.getAgentCorrectWordFileIds(id);
|
||||
agent.setCorrectWordFileIds(correctWordFileIds);
|
||||
agent.setCurrentVersionNo(agentSnapshotService.getCurrentVersionNo(id));
|
||||
|
||||
// 无需额外查询插件列表,已通过SQL查询出来
|
||||
return agent;
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -195,15 +197,35 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
|
||||
return super.insert(entity);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void deleteAgentByUserId(Long userId) {
|
||||
UpdateWrapper<AgentEntity> wrapper = new UpdateWrapper<>();
|
||||
wrapper.eq("user_id", userId);
|
||||
baseDao.delete(wrapper);
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<AgentDTO> getUserAgents(Long userId, String keyword, String searchType) {
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void deleteAgentByUserId(Long userId) {
|
||||
List<AgentEntity> agents = baseDao.selectList(new QueryWrapper<AgentEntity>()
|
||||
.select("id")
|
||||
.eq("user_id", userId));
|
||||
for (AgentEntity agent : agents) {
|
||||
deleteAgent(agent.getId());
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void deleteAgent(String agentId) {
|
||||
if (agentDao.selectByIdForUpdate(agentId) == null) {
|
||||
return;
|
||||
}
|
||||
deviceService.deleteByAgentId(agentId);
|
||||
agentChatHistoryService.deleteByAgentId(agentId, true, true);
|
||||
agentPluginMappingService.deleteByAgentId(agentId);
|
||||
agentContextProviderService.deleteByAgentId(agentId);
|
||||
correctWordFileService.deleteMappingsByAgentId(agentId);
|
||||
agentTagService.deleteAgentTags(agentId);
|
||||
agentSnapshotService.deleteByAgentId(agentId);
|
||||
deleteById(agentId);
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<AgentDTO> getUserAgents(Long userId, String keyword, String searchType) {
|
||||
QueryWrapper<AgentEntity> queryWrapper = new QueryWrapper<>();
|
||||
queryWrapper.eq("user_id", userId).orderByDesc("created_at");
|
||||
|
||||
@@ -314,36 +336,55 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean checkAgentPermission(String agentId, Long userId) {
|
||||
public boolean checkAgentPermission(String agentId, Long userId) {
|
||||
AgentEntity agent = agentDao.selectById(agentId);
|
||||
return hasAgentPermission(agent, userId);
|
||||
}
|
||||
|
||||
// 根据id更新智能体信息
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void updateAgentById(String agentId, AgentUpdateDTO dto) {
|
||||
AgentEntity existingEntity = getAgentEntityOrThrow(agentId);
|
||||
requireCurrentUserPermissionIfPresent(existingEntity);
|
||||
updateAgentEntity(agentId, dto, existingEntity);
|
||||
}
|
||||
|
||||
// 根据id更新智能体信息
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void updateAgentById(String agentId, AgentUpdateDTO dto) {
|
||||
updateAgentById(agentId, dto, true);
|
||||
}
|
||||
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void updateAgentById(String agentId, AgentUpdateDTO dto, Long userId) {
|
||||
AgentEntity existingEntity = getAgentEntityOrThrow(agentId);
|
||||
requireAgentPermission(existingEntity, userId);
|
||||
updateAgentEntity(agentId, dto, existingEntity);
|
||||
updateAgentById(agentId, dto, userId, true);
|
||||
}
|
||||
|
||||
private void updateAgentEntity(String agentId, AgentUpdateDTO dto, AgentEntity existingEntity) {
|
||||
// 只更新提供的非空字段
|
||||
if (dto.getAgentName() != null) {
|
||||
existingEntity.setAgentName(dto.getAgentName());
|
||||
}
|
||||
if (dto.getAgentCode() != null) {
|
||||
existingEntity.setAgentCode(dto.getAgentCode());
|
||||
}
|
||||
// 根据id更新智能体信息
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void updateAgentById(String agentId, AgentUpdateDTO dto, boolean createSnapshot) {
|
||||
updateAgentById(agentId, dto, null, createSnapshot);
|
||||
}
|
||||
|
||||
private void updateAgentById(String agentId, AgentUpdateDTO dto, Long userId, boolean createSnapshot) {
|
||||
AgentEntity lockedAgent = agentDao.selectByIdForUpdate(agentId);
|
||||
if (lockedAgent == null) {
|
||||
throw new RenException(ErrorCode.AGENT_NOT_FOUND);
|
||||
}
|
||||
if (userId == null) {
|
||||
requireCurrentUserPermissionIfPresent(lockedAgent);
|
||||
} else {
|
||||
requireAgentPermission(lockedAgent, userId);
|
||||
}
|
||||
|
||||
// 锁定后查询现有实体和关联配置
|
||||
AgentEntity existingEntity = this.getAgentById(agentId);
|
||||
if (createSnapshot && agentSnapshotService.getCurrentVersionNo(agentId) == 0) {
|
||||
agentSnapshotService.createSnapshot(agentId, "initial");
|
||||
}
|
||||
|
||||
// 只更新提供的非空字段
|
||||
if (dto.getAgentName() != null) {
|
||||
existingEntity.setAgentName(dto.getAgentName());
|
||||
}
|
||||
if (dto.getAgentCode() != null) {
|
||||
existingEntity.setAgentCode(dto.getAgentCode());
|
||||
}
|
||||
if (dto.getAsrModelId() != null) {
|
||||
existingEntity.setAsrModelId(dto.getAsrModelId());
|
||||
}
|
||||
@@ -481,16 +522,24 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
|
||||
}
|
||||
|
||||
// 更新替换词文件关联
|
||||
if (dto.getCorrectWordFileIds() != null) {
|
||||
correctWordFileService.saveAgentCorrectWords(agentId, dto.getCorrectWordFileIds());
|
||||
}
|
||||
|
||||
boolean b = validateLLMIntentParams(dto.getLlmModelId(), dto.getIntentModelId());
|
||||
if (!b) {
|
||||
throw new RenException(ErrorCode.LLM_INTENT_PARAMS_MISMATCH);
|
||||
}
|
||||
this.updateById(existingEntity);
|
||||
}
|
||||
if (dto.getCorrectWordFileIds() != null) {
|
||||
correctWordFileService.saveAgentCorrectWords(agentId, dto.getCorrectWordFileIds());
|
||||
}
|
||||
|
||||
// 更新智能体标签
|
||||
if (dto.getTagNames() != null || dto.getTagIds() != null) {
|
||||
agentTagService.saveAgentTags(agentId, dto.getTagIds(), dto.getTagNames());
|
||||
}
|
||||
|
||||
boolean b = validateLLMIntentParams(existingEntity.getLlmModelId(), existingEntity.getIntentModelId());
|
||||
if (!b) {
|
||||
throw new RenException(ErrorCode.LLM_INTENT_PARAMS_MISMATCH);
|
||||
}
|
||||
this.updateById(existingEntity);
|
||||
if (createSnapshot) {
|
||||
agentSnapshotService.createSnapshot(agentId, "config");
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
@@ -504,7 +553,7 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
|
||||
|
||||
AgentUpdateDTO agentUpdateDTO = new AgentUpdateDTO();
|
||||
agentUpdateDTO.setSummaryMemory(dto.getSummaryMemory());
|
||||
updateAgentById(device.getAgentId(), agentUpdateDTO, userId);
|
||||
updateAgentById(device.getAgentId(), agentUpdateDTO, userId, false);
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -512,19 +561,7 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
|
||||
public void deleteAgentById(String agentId, Long userId) {
|
||||
AgentEntity agent = getAgentEntityOrThrow(agentId);
|
||||
requireAgentPermission(agent, userId);
|
||||
|
||||
// 先删除关联的设备
|
||||
deviceService.deleteByAgentId(agentId);
|
||||
// 删除关联的聊天记录
|
||||
agentChatHistoryService.deleteByAgentId(agentId, true, true);
|
||||
// 删除关联的插件
|
||||
agentPluginMappingService.deleteByAgentId(agentId);
|
||||
// 删除关联的上下文源配置
|
||||
agentContextProviderService.deleteByAgentId(agentId);
|
||||
// 删除关联的替换词文件关联记录
|
||||
correctWordFileService.deleteMappingsByAgentId(agentId);
|
||||
// 再删除智能体
|
||||
baseDao.deleteById(agentId);
|
||||
deleteAgent(agentId);
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
+759
@@ -0,0 +1,759 @@
|
||||
package xiaozhi.modules.agent.service.impl;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collections;
|
||||
import java.util.Date;
|
||||
import java.util.HashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Objects;
|
||||
import java.util.TreeMap;
|
||||
import java.util.UUID;
|
||||
import java.util.function.Function;
|
||||
import java.util.stream.Collectors;
|
||||
|
||||
import org.apache.commons.lang3.StringUtils;
|
||||
import org.springframework.stereotype.Service;
|
||||
import org.springframework.transaction.annotation.Transactional;
|
||||
|
||||
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
|
||||
import com.baomidou.mybatisplus.core.metadata.IPage;
|
||||
import com.baomidou.mybatisplus.core.metadata.OrderItem;
|
||||
import com.baomidou.mybatisplus.extension.plugins.pagination.Page;
|
||||
import com.fasterxml.jackson.core.type.TypeReference;
|
||||
|
||||
import lombok.AllArgsConstructor;
|
||||
import xiaozhi.common.constant.Constant;
|
||||
import xiaozhi.common.exception.ErrorCode;
|
||||
import xiaozhi.common.exception.RenException;
|
||||
import xiaozhi.common.page.PageData;
|
||||
import xiaozhi.common.service.impl.BaseServiceImpl;
|
||||
import xiaozhi.common.utils.JsonUtils;
|
||||
import xiaozhi.modules.agent.dao.AgentDao;
|
||||
import xiaozhi.modules.agent.dao.AgentSnapshotDao;
|
||||
import xiaozhi.modules.agent.dao.AgentTagDao;
|
||||
import xiaozhi.modules.agent.dao.AgentTagRelationDao;
|
||||
import xiaozhi.modules.agent.dto.AgentSnapshotDataDTO;
|
||||
import xiaozhi.modules.agent.dto.AgentSnapshotPageDTO;
|
||||
import xiaozhi.modules.agent.dto.AgentSnapshotTagDTO;
|
||||
import xiaozhi.modules.agent.dto.AgentUpdateDTO;
|
||||
import xiaozhi.modules.agent.entity.AgentContextProviderEntity;
|
||||
import xiaozhi.modules.agent.entity.AgentEntity;
|
||||
import xiaozhi.modules.agent.entity.AgentPluginMapping;
|
||||
import xiaozhi.modules.agent.entity.AgentSnapshotEntity;
|
||||
import xiaozhi.modules.agent.entity.AgentTagEntity;
|
||||
import xiaozhi.modules.agent.entity.AgentTagRelationEntity;
|
||||
import xiaozhi.modules.agent.Enums.AgentSnapshotField;
|
||||
import xiaozhi.modules.agent.service.AgentChatHistoryService;
|
||||
import xiaozhi.modules.agent.service.AgentContextProviderService;
|
||||
import xiaozhi.modules.agent.service.AgentPluginMappingService;
|
||||
import xiaozhi.modules.agent.service.AgentSnapshotService;
|
||||
import xiaozhi.modules.agent.service.AgentTagService;
|
||||
import xiaozhi.modules.agent.vo.AgentInfoVO;
|
||||
import xiaozhi.modules.agent.vo.AgentSnapshotVO;
|
||||
import xiaozhi.modules.correctword.service.CorrectWordFileService;
|
||||
import xiaozhi.modules.model.entity.ModelConfigEntity;
|
||||
import xiaozhi.modules.model.service.ModelConfigService;
|
||||
import xiaozhi.modules.security.user.SecurityUser;
|
||||
|
||||
@Service
|
||||
@AllArgsConstructor
|
||||
public class AgentSnapshotServiceImpl extends BaseServiceImpl<AgentSnapshotDao, AgentSnapshotEntity>
|
||||
implements AgentSnapshotService {
|
||||
private static final int MAX_SNAPSHOTS_PER_AGENT = 100;
|
||||
private static final TypeReference<List<String>> STRING_LIST_TYPE = new TypeReference<>() {
|
||||
};
|
||||
private static final TypeReference<HashMap<String, Object>> PARAM_INFO_TYPE = new TypeReference<>() {
|
||||
};
|
||||
private static final TypeReference<Map<String, Object>> OBJECT_MAP_TYPE = new TypeReference<>() {
|
||||
};
|
||||
private static final String SECRET_PLACEHOLDER = "__SNAPSHOT_SECRET_REDACTED__";
|
||||
|
||||
private final AgentSnapshotDao agentSnapshotDao;
|
||||
private final AgentDao agentDao;
|
||||
private final AgentTagDao agentTagDao;
|
||||
private final AgentTagRelationDao agentTagRelationDao;
|
||||
private final AgentPluginMappingService agentPluginMappingService;
|
||||
private final AgentContextProviderService agentContextProviderService;
|
||||
private final AgentTagService agentTagService;
|
||||
private final AgentChatHistoryService agentChatHistoryService;
|
||||
private final ModelConfigService modelConfigService;
|
||||
private final CorrectWordFileService correctWordFileService;
|
||||
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void createSnapshot(String agentId, String source) {
|
||||
lockAgent(agentId);
|
||||
AgentInfoVO agent = getAgentInfo(agentId);
|
||||
AgentSnapshotDataDTO snapshotData = buildSnapshotData(agent);
|
||||
AgentSnapshotEntity previousSnapshot = agentSnapshotDao.selectLatestSnapshot(agentId);
|
||||
List<String> changedFields;
|
||||
if (previousSnapshot == null) {
|
||||
changedFields = List.of("initial");
|
||||
} else {
|
||||
AgentSnapshotDataDTO previousData = JsonUtils.parseObject(previousSnapshot.getSnapshotData(),
|
||||
AgentSnapshotDataDTO.class);
|
||||
changedFields = getChangedFields(previousData, snapshotData);
|
||||
}
|
||||
if (changedFields.isEmpty()) {
|
||||
return;
|
||||
}
|
||||
|
||||
insertSnapshot(agentId, agent.getUserId(), source, snapshotData, changedFields);
|
||||
}
|
||||
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public PageData<AgentSnapshotVO> page(String agentId, AgentSnapshotPageDTO params) {
|
||||
AgentSnapshotPageDTO pageParams = params == null ? new AgentSnapshotPageDTO() : params;
|
||||
Page<AgentSnapshotEntity> page = new Page<>(pageParams.pageOrDefault(), pageParams.limitOrDefault());
|
||||
page.addOrder(OrderItem.desc("version_no"));
|
||||
IPage<AgentSnapshotEntity> result = agentSnapshotDao.selectPage(page,
|
||||
new QueryWrapper<AgentSnapshotEntity>()
|
||||
.eq("agent_id", agentId)
|
||||
.le(pageParams.getMaxVersionNo() != null, "version_no", pageParams.getMaxVersionNo()));
|
||||
List<AgentSnapshotVO> list = result.getRecords().stream()
|
||||
.map(entity -> toVO(entity, false))
|
||||
.toList();
|
||||
return new PageData<>(list, result.getTotal());
|
||||
}
|
||||
|
||||
@Override
|
||||
public AgentSnapshotVO getSnapshot(String agentId, String snapshotId) {
|
||||
AgentSnapshotEntity entity = getSnapshotEntity(agentId, snapshotId);
|
||||
return toVO(entity, true);
|
||||
}
|
||||
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void restoreSnapshot(String agentId, String snapshotId) {
|
||||
AgentEntity agent = agentDao.selectByIdForUpdate(agentId);
|
||||
if (agent == null) {
|
||||
throw new RenException(ErrorCode.AGENT_NOT_FOUND);
|
||||
}
|
||||
|
||||
AgentSnapshotEntity snapshot = getSnapshotEntity(agentId, snapshotId);
|
||||
AgentSnapshotDataDTO data = JsonUtils.parseObject(snapshot.getSnapshotData(), AgentSnapshotDataDTO.class);
|
||||
if (data == null) {
|
||||
throw new RenException("快照数据为空,无法恢复");
|
||||
}
|
||||
|
||||
AgentInfoVO currentAgent = getAgentInfo(agentId);
|
||||
AgentSnapshotDataDTO currentData = buildSnapshotData(currentAgent);
|
||||
AgentSnapshotDataDTO restoreData = preserveCurrentSensitiveValues(data, currentData);
|
||||
List<String> changedFields = getChangedFields(currentData, restoreData);
|
||||
if (changedFields.isEmpty()) {
|
||||
return;
|
||||
}
|
||||
|
||||
applyAgentFields(agent, restoreData);
|
||||
validateRestoreParams(agent);
|
||||
applyMemoryPolicy(agent);
|
||||
agent.setUpdater(SecurityUser.getUserId());
|
||||
agent.setUpdatedAt(new Date());
|
||||
agentDao.updateById(agent);
|
||||
|
||||
restoreFunctions(agentId, restoreData.getFunctions());
|
||||
restoreContextProviders(agentId, restoreData.getContextProviders());
|
||||
correctWordFileService.saveAgentCorrectWords(agentId, nullToEmpty(restoreData.getCorrectWordFileIds()));
|
||||
restoreTags(agentId, restoreData);
|
||||
insertSnapshot(agentId, currentAgent.getUserId(), "restore", restoreData, changedFields,
|
||||
snapshot.getId(), snapshot.getVersionNo());
|
||||
}
|
||||
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void deleteSnapshot(String agentId, String snapshotId) {
|
||||
AgentSnapshotEntity entity = getSnapshotEntity(agentId, snapshotId);
|
||||
Integer maxVersionNo = agentSnapshotDao.selectMaxVersionNo(agentId);
|
||||
if (Objects.equals(entity.getVersionNo(), maxVersionNo)) {
|
||||
throw new RenException("最新历史版本不能删除");
|
||||
}
|
||||
deleteById(snapshotId);
|
||||
}
|
||||
|
||||
@Override
|
||||
public Integer getCurrentVersionNo(String agentId) {
|
||||
Integer maxVersionNo = agentSnapshotDao.selectMaxVersionNo(agentId);
|
||||
return maxVersionNo == null ? 0 : maxVersionNo;
|
||||
}
|
||||
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void deleteByAgentId(String agentId) {
|
||||
agentSnapshotDao.delete(new QueryWrapper<AgentSnapshotEntity>().eq("agent_id", agentId));
|
||||
}
|
||||
|
||||
private void insertSnapshot(String agentId, Long userId, String source, AgentSnapshotDataDTO snapshotData,
|
||||
List<String> changedFields) {
|
||||
insertSnapshot(agentId, userId, source, snapshotData, changedFields, null, null);
|
||||
}
|
||||
|
||||
private void insertSnapshot(String agentId, Long userId, String source, AgentSnapshotDataDTO snapshotData,
|
||||
List<String> changedFields, String restoreFromSnapshotId, Integer restoreFromVersionNo) {
|
||||
AgentSnapshotEntity entity = new AgentSnapshotEntity();
|
||||
entity.setId(UUID.randomUUID().toString().replace("-", ""));
|
||||
entity.setAgentId(agentId);
|
||||
entity.setUserId(userId);
|
||||
entity.setSnapshotData(JsonUtils.toJsonString(redactSnapshotData(snapshotData)));
|
||||
entity.setChangedFields(JsonUtils.toJsonString(changedFields));
|
||||
entity.setSource(StringUtils.defaultIfBlank(source, "config"));
|
||||
entity.setRestoreFromSnapshotId(restoreFromSnapshotId);
|
||||
entity.setRestoreFromVersionNo(restoreFromVersionNo);
|
||||
entity.setCreator(SecurityUser.getUserId());
|
||||
entity.setCreatedAt(new Date());
|
||||
int inserted = agentSnapshotDao.insertWithNextVersion(entity);
|
||||
if (inserted != 1) {
|
||||
throw new RenException("快照版本号生成失败");
|
||||
}
|
||||
pruneSnapshots(agentId);
|
||||
}
|
||||
|
||||
private void lockAgent(String agentId) {
|
||||
if (agentDao.selectByIdForUpdate(agentId) == null) {
|
||||
throw new RenException(ErrorCode.AGENT_NOT_FOUND);
|
||||
}
|
||||
}
|
||||
|
||||
private void pruneSnapshots(String agentId) {
|
||||
agentSnapshotDao.deleteOlderThanKeepLimit(agentId, MAX_SNAPSHOTS_PER_AGENT);
|
||||
}
|
||||
|
||||
private AgentSnapshotEntity getSnapshotEntity(String agentId, String snapshotId) {
|
||||
AgentSnapshotEntity entity = selectById(snapshotId);
|
||||
if (entity == null || !Objects.equals(agentId, entity.getAgentId())) {
|
||||
throw new RenException("快照不存在");
|
||||
}
|
||||
return entity;
|
||||
}
|
||||
|
||||
private AgentInfoVO getAgentInfo(String agentId) {
|
||||
AgentInfoVO agent = agentDao.selectAgentInfoById(agentId);
|
||||
if (agent == null) {
|
||||
throw new RenException(ErrorCode.AGENT_NOT_FOUND);
|
||||
}
|
||||
return agent;
|
||||
}
|
||||
|
||||
private AgentSnapshotDataDTO buildSnapshotData(String agentId) {
|
||||
return buildSnapshotData(getAgentInfo(agentId));
|
||||
}
|
||||
|
||||
private AgentSnapshotDataDTO buildSnapshotData(AgentInfoVO agent) {
|
||||
AgentSnapshotDataDTO data = new AgentSnapshotDataDTO();
|
||||
data.setAgentCode(agent.getAgentCode());
|
||||
data.setAgentName(agent.getAgentName());
|
||||
data.setAsrModelId(agent.getAsrModelId());
|
||||
data.setVadModelId(agent.getVadModelId());
|
||||
data.setLlmModelId(agent.getLlmModelId());
|
||||
data.setSlmModelId(agent.getSlmModelId());
|
||||
data.setVllmModelId(agent.getVllmModelId());
|
||||
data.setTtsModelId(agent.getTtsModelId());
|
||||
data.setTtsVoiceId(agent.getTtsVoiceId());
|
||||
data.setTtsLanguage(agent.getTtsLanguage());
|
||||
data.setTtsVolume(agent.getTtsVolume());
|
||||
data.setTtsRate(agent.getTtsRate());
|
||||
data.setTtsPitch(agent.getTtsPitch());
|
||||
data.setMemModelId(agent.getMemModelId());
|
||||
data.setIntentModelId(agent.getIntentModelId());
|
||||
data.setChatHistoryConf(agent.getChatHistoryConf());
|
||||
data.setSystemPrompt(agent.getSystemPrompt());
|
||||
data.setSummaryMemory(agent.getSummaryMemory());
|
||||
data.setLangCode(agent.getLangCode());
|
||||
data.setLanguage(agent.getLanguage());
|
||||
data.setSort(agent.getSort());
|
||||
data.setFunctions(toFunctionInfo(agent.getFunctions()));
|
||||
data.setContextProviders(getContextProviders(agent.getId()));
|
||||
data.setCorrectWordFileIds(nullToEmpty(correctWordFileService.getAgentCorrectWordFileIds(agent.getId())));
|
||||
List<AgentSnapshotTagDTO> tags = getSnapshotTags(agent.getId());
|
||||
data.setTags(tags);
|
||||
data.setTagNames(tags.stream().map(AgentSnapshotTagDTO::getTagName).toList());
|
||||
return data;
|
||||
}
|
||||
|
||||
private List<AgentSnapshotTagDTO> getSnapshotTags(String agentId) {
|
||||
List<AgentTagEntity> tags = agentTagDao.selectByAgentId(agentId);
|
||||
if (tags == null || tags.isEmpty()) {
|
||||
return Collections.emptyList();
|
||||
}
|
||||
List<AgentSnapshotTagDTO> snapshotTags = new ArrayList<>();
|
||||
for (int i = 0; i < tags.size(); i++) {
|
||||
AgentTagEntity tag = tags.get(i);
|
||||
AgentSnapshotTagDTO snapshotTag = new AgentSnapshotTagDTO();
|
||||
snapshotTag.setId(tag.getId());
|
||||
snapshotTag.setTagName(tag.getTagName());
|
||||
snapshotTag.setSort(i);
|
||||
snapshotTags.add(snapshotTag);
|
||||
}
|
||||
return snapshotTags;
|
||||
}
|
||||
|
||||
private List<AgentUpdateDTO.FunctionInfo> toFunctionInfo(List<AgentPluginMapping> mappings) {
|
||||
if (mappings == null || mappings.isEmpty()) {
|
||||
return Collections.emptyList();
|
||||
}
|
||||
return mappings.stream().map(mapping -> {
|
||||
AgentUpdateDTO.FunctionInfo info = new AgentUpdateDTO.FunctionInfo();
|
||||
info.setPluginId(mapping.getPluginId());
|
||||
info.setParamInfo(parseParamInfo(mapping.getParamInfo()));
|
||||
return info;
|
||||
}).toList();
|
||||
}
|
||||
|
||||
private HashMap<String, Object> parseParamInfo(String paramInfo) {
|
||||
if (StringUtils.isBlank(paramInfo)) {
|
||||
return new HashMap<>();
|
||||
}
|
||||
return JsonUtils.parseObject(paramInfo, PARAM_INFO_TYPE);
|
||||
}
|
||||
|
||||
private List<xiaozhi.modules.agent.dto.ContextProviderDTO> getContextProviders(String agentId) {
|
||||
AgentContextProviderEntity contextEntity = agentContextProviderService.getByAgentId(agentId);
|
||||
if (contextEntity == null || contextEntity.getContextProviders() == null) {
|
||||
return Collections.emptyList();
|
||||
}
|
||||
return contextEntity.getContextProviders();
|
||||
}
|
||||
|
||||
private List<String> getChangedFields(AgentSnapshotDataDTO current, AgentUpdateDTO pendingUpdate) {
|
||||
if (pendingUpdate == null) {
|
||||
return List.of("restore");
|
||||
}
|
||||
|
||||
List<String> fields = new ArrayList<>();
|
||||
for (AgentSnapshotField field : AgentSnapshotField.values()) {
|
||||
Object nextValue = field.updateValue(pendingUpdate);
|
||||
if (nextValue != null && isChanged(field, field.snapshotValue(current), nextValue)) {
|
||||
fields.add(field.getFieldName());
|
||||
}
|
||||
}
|
||||
return fields;
|
||||
}
|
||||
|
||||
private List<String> getChangedFields(AgentSnapshotDataDTO current, AgentSnapshotDataDTO next) {
|
||||
if (next == null) {
|
||||
return List.of("restore");
|
||||
}
|
||||
if (current == null) {
|
||||
current = new AgentSnapshotDataDTO();
|
||||
}
|
||||
|
||||
List<String> fields = new ArrayList<>();
|
||||
for (AgentSnapshotField field : AgentSnapshotField.values()) {
|
||||
if (isChanged(field, field.snapshotValue(current), field.snapshotValue(next))) {
|
||||
fields.add(field.getFieldName());
|
||||
}
|
||||
}
|
||||
return fields;
|
||||
}
|
||||
|
||||
private boolean isChanged(AgentSnapshotField field, Object current, Object next) {
|
||||
return !Objects.equals(normalizeForCompare(field, current), normalizeForCompare(field, next));
|
||||
}
|
||||
|
||||
@SuppressWarnings("unchecked")
|
||||
private Object normalizeForCompare(AgentSnapshotField field, Object value) {
|
||||
return switch (field) {
|
||||
case FUNCTIONS -> normalizeFunctions((List<AgentUpdateDTO.FunctionInfo>) value);
|
||||
case CONTEXT_PROVIDERS -> normalizeSortedJsonList((List<?>) value);
|
||||
case CORRECT_WORD_FILE_IDS -> normalizeStringList((List<String>) value);
|
||||
case TAG_NAMES -> normalizeStringList((List<String>) value);
|
||||
default -> value;
|
||||
};
|
||||
}
|
||||
|
||||
private Map<String, Object> normalizeFunctions(List<AgentUpdateDTO.FunctionInfo> functions) {
|
||||
Map<String, Object> result = new TreeMap<>();
|
||||
if (functions == null) {
|
||||
return result;
|
||||
}
|
||||
for (AgentUpdateDTO.FunctionInfo function : functions) {
|
||||
if (function == null || StringUtils.isBlank(function.getPluginId())) {
|
||||
continue;
|
||||
}
|
||||
result.put(function.getPluginId(), normalizeMap(function.getParamInfo()));
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
private List<String> normalizeStringList(List<String> values) {
|
||||
if (values == null) {
|
||||
return Collections.emptyList();
|
||||
}
|
||||
return values.stream()
|
||||
.filter(StringUtils::isNotBlank)
|
||||
.sorted()
|
||||
.toList();
|
||||
}
|
||||
|
||||
private <T> List<String> normalizeSortedJsonList(List<T> values) {
|
||||
if (values == null) {
|
||||
return Collections.emptyList();
|
||||
}
|
||||
return values.stream()
|
||||
.filter(Objects::nonNull)
|
||||
.map(value -> JsonUtils.toJsonString(normalizeValue(value)))
|
||||
.sorted()
|
||||
.toList();
|
||||
}
|
||||
|
||||
private Object normalizeValue(Object value) {
|
||||
if (value == null) {
|
||||
return null;
|
||||
}
|
||||
if (value instanceof Map<?, ?> map) {
|
||||
return normalizeMap(map);
|
||||
}
|
||||
if (value instanceof List<?> list) {
|
||||
return list.stream().map(this::normalizeValue).toList();
|
||||
}
|
||||
if (isScalarValue(value)) {
|
||||
return value;
|
||||
}
|
||||
return normalizeMap(JsonUtils.parseObject(JsonUtils.toJsonString(value), OBJECT_MAP_TYPE));
|
||||
}
|
||||
|
||||
private boolean isScalarValue(Object value) {
|
||||
return value instanceof CharSequence
|
||||
|| value instanceof Number
|
||||
|| value instanceof Boolean
|
||||
|| value instanceof Character
|
||||
|| value instanceof Enum<?>
|
||||
|| value instanceof Date;
|
||||
}
|
||||
|
||||
private Map<String, Object> normalizeMap(Map<?, ?> map) {
|
||||
Map<String, Object> result = new TreeMap<>();
|
||||
if (map == null) {
|
||||
return result;
|
||||
}
|
||||
map.forEach((key, value) -> {
|
||||
if (key != null) {
|
||||
String keyText = String.valueOf(key);
|
||||
result.put(keyText, isSensitiveKey(keyText) ? SECRET_PLACEHOLDER : normalizeValue(value));
|
||||
}
|
||||
});
|
||||
return result;
|
||||
}
|
||||
|
||||
private void applyAgentFields(AgentEntity agent, AgentSnapshotDataDTO data) {
|
||||
for (AgentSnapshotField field : AgentSnapshotField.values()) {
|
||||
field.applyTo(agent, data);
|
||||
}
|
||||
}
|
||||
|
||||
private void validateRestoreParams(AgentEntity agent) {
|
||||
if (agent == null || StringUtils.isBlank(agent.getLlmModelId())) {
|
||||
return;
|
||||
}
|
||||
|
||||
ModelConfigEntity llmModelData = modelConfigService.selectById(agent.getLlmModelId());
|
||||
if (llmModelData == null || llmModelData.getConfigJson() == null) {
|
||||
throw new RenException(ErrorCode.LLM_INTENT_PARAMS_MISMATCH);
|
||||
}
|
||||
Object typeValue = llmModelData.getConfigJson().get("type");
|
||||
String type = typeValue == null ? "" : typeValue.toString();
|
||||
if ("openai".equals(type) || "ollama".equals(type)) {
|
||||
return;
|
||||
}
|
||||
if ("Intent_function_call".equals(agent.getIntentModelId())) {
|
||||
throw new RenException(ErrorCode.LLM_INTENT_PARAMS_MISMATCH);
|
||||
}
|
||||
}
|
||||
|
||||
private void applyMemoryPolicy(AgentEntity agent) {
|
||||
if (agent == null || StringUtils.isBlank(agent.getMemModelId())) {
|
||||
return;
|
||||
}
|
||||
if (Constant.MEMORY_NO_MEM.equals(agent.getMemModelId())) {
|
||||
agentChatHistoryService.deleteByAgentId(agent.getId(), true, true);
|
||||
agent.setSummaryMemory("");
|
||||
} else if (Constant.MEMORY_MEM_REPORT_ONLY.equals(agent.getMemModelId())) {
|
||||
agent.setSummaryMemory("");
|
||||
}
|
||||
}
|
||||
|
||||
private void restoreFunctions(String agentId, List<AgentUpdateDTO.FunctionInfo> functions) {
|
||||
agentPluginMappingService.deleteByAgentId(agentId);
|
||||
if (functions == null || functions.isEmpty()) {
|
||||
return;
|
||||
}
|
||||
List<AgentPluginMapping> mappings = functions.stream().map(info -> {
|
||||
AgentPluginMapping mapping = new AgentPluginMapping();
|
||||
mapping.setAgentId(agentId);
|
||||
mapping.setPluginId(info.getPluginId());
|
||||
mapping.setParamInfo(JsonUtils.toJsonString(info.getParamInfo()));
|
||||
return mapping;
|
||||
}).toList();
|
||||
agentPluginMappingService.saveBatch(mappings);
|
||||
}
|
||||
|
||||
private void restoreContextProviders(String agentId,
|
||||
List<xiaozhi.modules.agent.dto.ContextProviderDTO> contextProviders) {
|
||||
AgentContextProviderEntity entity = new AgentContextProviderEntity();
|
||||
entity.setAgentId(agentId);
|
||||
entity.setContextProviders(nullToEmpty(contextProviders));
|
||||
entity.setUpdater(SecurityUser.getUserId());
|
||||
entity.setUpdatedAt(new Date());
|
||||
agentContextProviderService.saveOrUpdateByAgentId(entity);
|
||||
}
|
||||
|
||||
private void restoreTags(String agentId, AgentSnapshotDataDTO data) {
|
||||
List<AgentSnapshotTagDTO> tags = data.getTags();
|
||||
if (tags == null || tags.isEmpty()) {
|
||||
agentTagService.saveAgentTags(agentId, null, nullToEmpty(data.getTagNames()));
|
||||
return;
|
||||
}
|
||||
|
||||
agentTagRelationDao.deleteByAgentId(agentId);
|
||||
List<AgentTagRelationEntity> relations = new ArrayList<>();
|
||||
Date now = new Date();
|
||||
for (int i = 0; i < tags.size(); i++) {
|
||||
AgentSnapshotTagDTO snapshotTag = tags.get(i);
|
||||
String tagId = restoreTag(snapshotTag, now);
|
||||
if (StringUtils.isBlank(tagId)) {
|
||||
continue;
|
||||
}
|
||||
AgentTagRelationEntity relation = new AgentTagRelationEntity();
|
||||
relation.setId(UUID.randomUUID().toString().replace("-", ""));
|
||||
relation.setAgentId(agentId);
|
||||
relation.setTagId(tagId);
|
||||
relation.setSort(i);
|
||||
relation.setCreator(SecurityUser.getUserId());
|
||||
relation.setCreatedAt(now);
|
||||
relation.setUpdater(SecurityUser.getUserId());
|
||||
relation.setUpdatedAt(now);
|
||||
relations.add(relation);
|
||||
}
|
||||
if (!relations.isEmpty()) {
|
||||
agentTagRelationDao.batchInsertRelation(relations);
|
||||
}
|
||||
}
|
||||
|
||||
private String restoreTag(AgentSnapshotTagDTO snapshotTag, Date now) {
|
||||
if (snapshotTag == null || StringUtils.isBlank(snapshotTag.getTagName())) {
|
||||
return null;
|
||||
}
|
||||
|
||||
AgentTagEntity tag = null;
|
||||
if (StringUtils.isNotBlank(snapshotTag.getId())) {
|
||||
tag = agentTagDao.selectById(snapshotTag.getId());
|
||||
if (tag != null && Objects.equals(tag.getDeleted(), 1)) {
|
||||
tag = null;
|
||||
}
|
||||
}
|
||||
if (tag == null) {
|
||||
tag = agentTagDao.selectOne(new QueryWrapper<AgentTagEntity>()
|
||||
.eq("tag_name", snapshotTag.getTagName())
|
||||
.eq("deleted", 0)
|
||||
.last("LIMIT 1"));
|
||||
}
|
||||
if (tag != null) {
|
||||
return tag.getId();
|
||||
}
|
||||
|
||||
AgentTagEntity newTag = new AgentTagEntity();
|
||||
newTag.setId(UUID.randomUUID().toString().replace("-", ""));
|
||||
newTag.setTagName(snapshotTag.getTagName());
|
||||
newTag.setSort(snapshotTag.getSort() == null ? 0 : snapshotTag.getSort());
|
||||
newTag.setDeleted(0);
|
||||
newTag.setCreator(SecurityUser.getUserId());
|
||||
newTag.setCreatedAt(now);
|
||||
newTag.setUpdater(SecurityUser.getUserId());
|
||||
newTag.setUpdatedAt(now);
|
||||
agentTagDao.insert(newTag);
|
||||
return newTag.getId();
|
||||
}
|
||||
|
||||
private AgentSnapshotVO toVO(AgentSnapshotEntity entity, boolean includeData) {
|
||||
AgentSnapshotVO vo = new AgentSnapshotVO();
|
||||
vo.setId(entity.getId());
|
||||
vo.setAgentId(entity.getAgentId());
|
||||
vo.setUserId(entity.getUserId());
|
||||
vo.setVersionNo(entity.getVersionNo());
|
||||
vo.setChangedFields(parseChangedFields(entity.getChangedFields()));
|
||||
vo.setFieldOrder(AgentSnapshotField.names());
|
||||
vo.setSource(entity.getSource());
|
||||
vo.setRestoreFromSnapshotId(entity.getRestoreFromSnapshotId());
|
||||
vo.setRestoreFromVersionNo(entity.getRestoreFromVersionNo());
|
||||
vo.setCreator(entity.getCreator());
|
||||
vo.setCreatedAt(entity.getCreatedAt());
|
||||
if (includeData) {
|
||||
vo.setSnapshotData(redactSnapshotData(JsonUtils.parseObject(entity.getSnapshotData(),
|
||||
AgentSnapshotDataDTO.class)));
|
||||
vo.setAfterSnapshotData(redactSnapshotData(getAfterSnapshotData(entity)));
|
||||
}
|
||||
return vo;
|
||||
}
|
||||
|
||||
private AgentSnapshotDataDTO getAfterSnapshotData(AgentSnapshotEntity entity) {
|
||||
AgentSnapshotEntity nextSnapshot = agentSnapshotDao.selectNextSnapshot(entity.getAgentId(), entity.getVersionNo());
|
||||
if (nextSnapshot != null) {
|
||||
return JsonUtils.parseObject(nextSnapshot.getSnapshotData(), AgentSnapshotDataDTO.class);
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
private List<String> parseChangedFields(String changedFields) {
|
||||
if (StringUtils.isBlank(changedFields)) {
|
||||
return Collections.emptyList();
|
||||
}
|
||||
return JsonUtils.parseObject(changedFields, STRING_LIST_TYPE);
|
||||
}
|
||||
|
||||
private AgentSnapshotDataDTO redactSnapshotData(AgentSnapshotDataDTO data) {
|
||||
if (data == null) {
|
||||
return null;
|
||||
}
|
||||
AgentSnapshotDataDTO copy = JsonUtils.parseObject(JsonUtils.toJsonString(data), AgentSnapshotDataDTO.class);
|
||||
if (copy.getFunctions() != null) {
|
||||
copy.getFunctions().forEach(function -> {
|
||||
if (function != null) {
|
||||
function.setParamInfo(redactSensitiveMap(function.getParamInfo()));
|
||||
}
|
||||
});
|
||||
}
|
||||
if (copy.getContextProviders() != null) {
|
||||
copy.getContextProviders().forEach(provider -> {
|
||||
if (provider != null) {
|
||||
provider.setHeaders(redactSensitiveMap(provider.getHeaders()));
|
||||
}
|
||||
});
|
||||
}
|
||||
return copy;
|
||||
}
|
||||
|
||||
private AgentSnapshotDataDTO preserveCurrentSensitiveValues(AgentSnapshotDataDTO target, AgentSnapshotDataDTO current) {
|
||||
if (target == null) {
|
||||
return null;
|
||||
}
|
||||
AgentSnapshotDataDTO copy = JsonUtils.parseObject(JsonUtils.toJsonString(target), AgentSnapshotDataDTO.class);
|
||||
Map<String, AgentUpdateDTO.FunctionInfo> currentFunctions = nullToEmpty(current == null ? null : current.getFunctions())
|
||||
.stream()
|
||||
.filter(function -> function != null && StringUtils.isNotBlank(function.getPluginId()))
|
||||
.collect(Collectors.toMap(AgentUpdateDTO.FunctionInfo::getPluginId, Function.identity(),
|
||||
(left, right) -> left));
|
||||
if (copy.getFunctions() != null) {
|
||||
copy.getFunctions().forEach(function -> {
|
||||
if (function == null) {
|
||||
return;
|
||||
}
|
||||
AgentUpdateDTO.FunctionInfo currentFunction = currentFunctions.get(function.getPluginId());
|
||||
function.setParamInfo(preserveCurrentSensitiveMap(function.getParamInfo(),
|
||||
currentFunction == null ? null : currentFunction.getParamInfo()));
|
||||
});
|
||||
}
|
||||
|
||||
Map<String, xiaozhi.modules.agent.dto.ContextProviderDTO> currentProviders = nullToEmpty(
|
||||
current == null ? null : current.getContextProviders())
|
||||
.stream()
|
||||
.filter(provider -> provider != null && StringUtils.isNotBlank(provider.getUrl()))
|
||||
.collect(Collectors.toMap(xiaozhi.modules.agent.dto.ContextProviderDTO::getUrl, Function.identity(),
|
||||
(left, right) -> left));
|
||||
if (copy.getContextProviders() != null) {
|
||||
copy.getContextProviders().forEach(provider -> {
|
||||
if (provider == null) {
|
||||
return;
|
||||
}
|
||||
xiaozhi.modules.agent.dto.ContextProviderDTO currentProvider = currentProviders.get(provider.getUrl());
|
||||
provider.setHeaders(preserveCurrentSensitiveMap(provider.getHeaders(),
|
||||
currentProvider == null ? null : currentProvider.getHeaders()));
|
||||
});
|
||||
}
|
||||
return copy;
|
||||
}
|
||||
|
||||
private HashMap<String, Object> redactSensitiveMap(Map<?, ?> map) {
|
||||
HashMap<String, Object> result = new HashMap<>();
|
||||
if (map == null) {
|
||||
return result;
|
||||
}
|
||||
map.forEach((key, value) -> {
|
||||
if (key != null) {
|
||||
String keyText = String.valueOf(key);
|
||||
result.put(keyText, isSensitiveKey(keyText) ? SECRET_PLACEHOLDER : redactSensitiveValue(value));
|
||||
}
|
||||
});
|
||||
return result;
|
||||
}
|
||||
|
||||
private Object redactSensitiveValue(Object value) {
|
||||
if (value instanceof Map<?, ?> map) {
|
||||
return redactSensitiveMap(map);
|
||||
}
|
||||
if (value instanceof List<?> list) {
|
||||
return list.stream().map(this::redactSensitiveValue).toList();
|
||||
}
|
||||
return value;
|
||||
}
|
||||
|
||||
private HashMap<String, Object> preserveCurrentSensitiveMap(Map<?, ?> target, Map<?, ?> current) {
|
||||
HashMap<String, Object> result = new HashMap<>();
|
||||
if (target == null) {
|
||||
return result;
|
||||
}
|
||||
target.forEach((key, value) -> {
|
||||
if (key == null) {
|
||||
return;
|
||||
}
|
||||
String keyText = String.valueOf(key);
|
||||
Object currentValue = getMapValue(current, keyText);
|
||||
if (isSensitiveKey(keyText)) {
|
||||
result.put(keyText, currentValue == null || SECRET_PLACEHOLDER.equals(currentValue) ? "" : currentValue);
|
||||
} else {
|
||||
result.put(keyText, preserveCurrentSensitiveValue(value, currentValue));
|
||||
}
|
||||
});
|
||||
return result;
|
||||
}
|
||||
|
||||
private Object preserveCurrentSensitiveValue(Object target, Object current) {
|
||||
if (target instanceof Map<?, ?> targetMap) {
|
||||
return preserveCurrentSensitiveMap(targetMap, current instanceof Map<?, ?> currentMap ? currentMap : null);
|
||||
}
|
||||
if (target instanceof List<?> targetList) {
|
||||
List<?> currentList = current instanceof List<?> list ? list : Collections.emptyList();
|
||||
List<Object> result = new ArrayList<>();
|
||||
for (int i = 0; i < targetList.size(); i++) {
|
||||
Object currentItem = i < currentList.size() ? currentList.get(i) : null;
|
||||
result.add(preserveCurrentSensitiveValue(targetList.get(i), currentItem));
|
||||
}
|
||||
return result;
|
||||
}
|
||||
return target;
|
||||
}
|
||||
|
||||
private Object getMapValue(Map<?, ?> map, String key) {
|
||||
if (map == null) {
|
||||
return null;
|
||||
}
|
||||
if (map.containsKey(key)) {
|
||||
return map.get(key);
|
||||
}
|
||||
for (Map.Entry<?, ?> entry : map.entrySet()) {
|
||||
if (entry.getKey() != null && Objects.equals(key, String.valueOf(entry.getKey()))) {
|
||||
return entry.getValue();
|
||||
}
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
private boolean isSensitiveKey(String key) {
|
||||
String normalized = StringUtils.defaultString(key).toLowerCase().replaceAll("[^a-z0-9]", "");
|
||||
return normalized.equals("authorization")
|
||||
|| normalized.equals("token")
|
||||
|| normalized.endsWith("token")
|
||||
|| normalized.contains("apikey")
|
||||
|| normalized.contains("appkey")
|
||||
|| normalized.contains("accesskey")
|
||||
|| normalized.contains("privatekey")
|
||||
|| normalized.contains("password")
|
||||
|| normalized.contains("passwd")
|
||||
|| normalized.contains("secret")
|
||||
|| normalized.contains("credential");
|
||||
}
|
||||
|
||||
private <T> List<T> nullToEmpty(List<T> list) {
|
||||
return list == null ? Collections.emptyList() : list;
|
||||
}
|
||||
}
|
||||
+31
@@ -0,0 +1,31 @@
|
||||
package xiaozhi.modules.agent.typehandler;
|
||||
|
||||
import java.util.Collections;
|
||||
import java.util.List;
|
||||
|
||||
import org.apache.commons.lang3.StringUtils;
|
||||
|
||||
import com.baomidou.mybatisplus.extension.handlers.AbstractJsonTypeHandler;
|
||||
import com.fasterxml.jackson.core.type.TypeReference;
|
||||
|
||||
import xiaozhi.common.utils.JsonUtils;
|
||||
import xiaozhi.modules.agent.dto.ContextProviderDTO;
|
||||
|
||||
public class ContextProviderListTypeHandler extends AbstractJsonTypeHandler<List<ContextProviderDTO>> {
|
||||
private static final TypeReference<List<ContextProviderDTO>> CONTEXT_PROVIDER_LIST_TYPE = new TypeReference<>() {
|
||||
};
|
||||
|
||||
@Override
|
||||
protected List<ContextProviderDTO> parse(String json) {
|
||||
if (StringUtils.isBlank(json)) {
|
||||
return Collections.emptyList();
|
||||
}
|
||||
List<ContextProviderDTO> providers = JsonUtils.parseObject(json, CONTEXT_PROVIDER_LIST_TYPE);
|
||||
return providers == null ? Collections.emptyList() : providers;
|
||||
}
|
||||
|
||||
@Override
|
||||
protected String toJson(List<ContextProviderDTO> obj) {
|
||||
return JsonUtils.toJsonString(obj == null ? Collections.emptyList() : obj);
|
||||
}
|
||||
}
|
||||
@@ -1,7 +1,5 @@
|
||||
package xiaozhi.modules.agent.vo;
|
||||
|
||||
import com.baomidou.mybatisplus.annotation.TableField;
|
||||
import com.baomidou.mybatisplus.extension.handlers.JacksonTypeHandler;
|
||||
import io.swagger.v3.oas.annotations.media.Schema;
|
||||
import lombok.Data;
|
||||
import lombok.EqualsAndHashCode;
|
||||
@@ -20,7 +18,6 @@ import java.util.List;
|
||||
public class AgentInfoVO extends AgentEntity
|
||||
{
|
||||
@Schema(description = "插件列表Id")
|
||||
@TableField(typeHandler = JacksonTypeHandler.class)
|
||||
private List<AgentPluginMapping> functions;
|
||||
|
||||
@Schema(description = "上下文源配置")
|
||||
@@ -28,4 +25,7 @@ public class AgentInfoVO extends AgentEntity
|
||||
|
||||
@Schema(description = "替换词文件ID列表")
|
||||
private List<String> correctWordFileIds;
|
||||
|
||||
@Schema(description = "当前配置版本号")
|
||||
private Integer currentVersionNo;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
package xiaozhi.modules.agent.vo;
|
||||
|
||||
import java.util.Date;
|
||||
import java.util.List;
|
||||
|
||||
import io.swagger.v3.oas.annotations.media.Schema;
|
||||
import lombok.Data;
|
||||
import xiaozhi.modules.agent.dto.AgentSnapshotDataDTO;
|
||||
|
||||
@Data
|
||||
@Schema(description = "智能体配置快照")
|
||||
public class AgentSnapshotVO {
|
||||
private String id;
|
||||
private String agentId;
|
||||
@Schema(description = "所属用户ID,表示该快照归属的智能体所有者")
|
||||
private Long userId;
|
||||
private Integer versionNo;
|
||||
private List<String> changedFields;
|
||||
private List<String> fieldOrder;
|
||||
private String source;
|
||||
@Schema(description = "恢复来源快照ID,仅恢复前自动备份版本有值")
|
||||
private String restoreFromSnapshotId;
|
||||
@Schema(description = "恢复来源版本号,仅恢复前自动备份版本有值")
|
||||
private Integer restoreFromVersionNo;
|
||||
@Schema(description = "创建者,表示触发本次快照写入的操作人")
|
||||
private Long creator;
|
||||
private Date createdAt;
|
||||
private AgentSnapshotDataDTO snapshotData;
|
||||
private AgentSnapshotDataDTO afterSnapshotData;
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
-- liquibase formatted sql
|
||||
|
||||
-- changeset tykechen:202607071530
|
||||
CREATE TABLE IF NOT EXISTS `ai_agent_snapshot` (
|
||||
`id` VARCHAR(32) NOT NULL COMMENT '快照ID',
|
||||
`agent_id` VARCHAR(32) NOT NULL COMMENT '智能体ID',
|
||||
`user_id` BIGINT DEFAULT NULL COMMENT '所属用户ID',
|
||||
`version_no` INT UNSIGNED NOT NULL COMMENT '版本号',
|
||||
`snapshot_data` JSON NOT NULL DEFAULT (JSON_OBJECT()) COMMENT '快照数据',
|
||||
`changed_fields` JSON DEFAULT NULL COMMENT '变更字段',
|
||||
`source` VARCHAR(32) DEFAULT 'config' COMMENT '快照来源',
|
||||
`restore_from_snapshot_id` VARCHAR(32) DEFAULT NULL COMMENT '恢复来源快照ID',
|
||||
`restore_from_version_no` INT UNSIGNED DEFAULT NULL COMMENT '恢复来源版本号',
|
||||
`creator` BIGINT DEFAULT NULL COMMENT '创建者',
|
||||
`created_at` DATETIME DEFAULT CURRENT_TIMESTAMP COMMENT '创建时间',
|
||||
PRIMARY KEY (`id`),
|
||||
UNIQUE KEY `uk_agent_version` (`agent_id`, `version_no`),
|
||||
INDEX `idx_agent_created_at` (`agent_id`, `created_at`),
|
||||
INDEX `idx_snapshot_user_created_at` (`user_id`, `created_at`)
|
||||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='智能体配置快照表';
|
||||
@@ -697,3 +697,10 @@ databaseChangeLog:
|
||||
- sqlFile:
|
||||
encoding: utf8
|
||||
path: classpath:db/changelog/202607011405.sql
|
||||
- changeSet:
|
||||
id: 202607071530
|
||||
author: tykechen
|
||||
changes:
|
||||
- sqlFile:
|
||||
encoding: utf8
|
||||
path: classpath:db/changelog/202607071530.sql
|
||||
|
||||
@@ -25,9 +25,6 @@
|
||||
<result column="memModelId" property="memModelId"/>
|
||||
<result column="intentModelId" property="intentModelId"/>
|
||||
|
||||
<result column="functions" property="functions"
|
||||
typeHandler="com.baomidou.mybatisplus.extension.handlers.JacksonTypeHandler"/>
|
||||
|
||||
<result column="chatHistoryConf" property="chatHistoryConf"/>
|
||||
<result column="systemPrompt" property="systemPrompt"/>
|
||||
<result column="summaryMemory" property="summaryMemory"/>
|
||||
@@ -38,6 +35,13 @@
|
||||
<result column="createdAt" property="createdAt"/>
|
||||
<result column="updater" property="updater"/>
|
||||
<result column="updatedAt" property="updatedAt"/>
|
||||
<collection property="functions"
|
||||
ofType="xiaozhi.modules.agent.entity.AgentPluginMapping">
|
||||
<id column="functionId" property="id"/>
|
||||
<result column="functionAgentId" property="agentId"/>
|
||||
<result column="pluginId" property="pluginId"/>
|
||||
<result column="paramInfo" property="paramInfo"/>
|
||||
</collection>
|
||||
</resultMap>
|
||||
|
||||
<select id="selectAgentInfoById" resultMap="AgentInfoMap">
|
||||
@@ -58,19 +62,6 @@
|
||||
a.tts_pitch AS ttsPitch,
|
||||
a.mem_model_id AS memModelId,
|
||||
a.intent_model_id AS intentModelId,
|
||||
COALESCE(
|
||||
(SELECT JSON_ARRAYAGG(
|
||||
JSON_OBJECT(
|
||||
'id', m.id,
|
||||
'agentId', m.agent_id,
|
||||
'pluginId', m.plugin_id,
|
||||
'paramInfo', m.param_info
|
||||
)
|
||||
)
|
||||
FROM ai_agent_plugin_mapping m
|
||||
WHERE m.agent_id = a.id),
|
||||
JSON_ARRAY()
|
||||
) AS functions,
|
||||
a.chat_history_conf AS chatHistoryConf,
|
||||
a.system_prompt AS systemPrompt,
|
||||
a.summary_memory AS summaryMemory,
|
||||
@@ -80,8 +71,22 @@
|
||||
a.creator,
|
||||
a.created_at AS createdAt,
|
||||
a.updater,
|
||||
a.updated_at AS updatedAt
|
||||
a.updated_at AS updatedAt,
|
||||
f.id AS functionId,
|
||||
f.agent_id AS functionAgentId,
|
||||
f.plugin_id AS pluginId,
|
||||
f.param_info AS paramInfo
|
||||
FROM ai_agent a
|
||||
LEFT JOIN ai_agent_plugin_mapping f ON f.agent_id = a.id
|
||||
WHERE a.id = #{agentId}
|
||||
ORDER BY f.id ASC
|
||||
</select>
|
||||
</mapper>
|
||||
|
||||
<select id="selectByIdForUpdate" resultType="xiaozhi.modules.agent.entity.AgentEntity">
|
||||
SELECT *
|
||||
FROM ai_agent
|
||||
WHERE id = #{agentId}
|
||||
FOR UPDATE
|
||||
</select>
|
||||
|
||||
</mapper>
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<!DOCTYPE mapper PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" "http://mybatis.org/dtd/mybatis-3-mapper.dtd">
|
||||
<mapper namespace="xiaozhi.modules.agent.dao.AgentSnapshotDao">
|
||||
|
||||
<select id="selectMaxVersionNo" resultType="java.lang.Integer">
|
||||
SELECT COALESCE(MAX(version_no), 0)
|
||||
FROM ai_agent_snapshot
|
||||
WHERE agent_id = #{agentId}
|
||||
</select>
|
||||
|
||||
<select id="selectLatestSnapshot" resultType="xiaozhi.modules.agent.entity.AgentSnapshotEntity">
|
||||
SELECT id,
|
||||
agent_id AS agentId,
|
||||
user_id AS userId,
|
||||
version_no AS versionNo,
|
||||
snapshot_data AS snapshotData,
|
||||
changed_fields AS changedFields,
|
||||
source,
|
||||
restore_from_snapshot_id AS restoreFromSnapshotId,
|
||||
restore_from_version_no AS restoreFromVersionNo,
|
||||
creator,
|
||||
created_at AS createdAt
|
||||
FROM ai_agent_snapshot
|
||||
WHERE agent_id = #{agentId}
|
||||
ORDER BY version_no DESC
|
||||
LIMIT 1
|
||||
</select>
|
||||
|
||||
<select id="selectNextSnapshot" resultType="xiaozhi.modules.agent.entity.AgentSnapshotEntity">
|
||||
SELECT id,
|
||||
agent_id AS agentId,
|
||||
user_id AS userId,
|
||||
version_no AS versionNo,
|
||||
snapshot_data AS snapshotData,
|
||||
changed_fields AS changedFields,
|
||||
source,
|
||||
restore_from_snapshot_id AS restoreFromSnapshotId,
|
||||
restore_from_version_no AS restoreFromVersionNo,
|
||||
creator,
|
||||
created_at AS createdAt
|
||||
FROM ai_agent_snapshot
|
||||
WHERE agent_id = #{agentId}
|
||||
AND version_no > #{versionNo}
|
||||
ORDER BY version_no ASC
|
||||
LIMIT 1
|
||||
</select>
|
||||
|
||||
<insert id="insertWithNextVersion">
|
||||
INSERT INTO ai_agent_snapshot (
|
||||
id,
|
||||
agent_id,
|
||||
user_id,
|
||||
version_no,
|
||||
snapshot_data,
|
||||
changed_fields,
|
||||
source,
|
||||
restore_from_snapshot_id,
|
||||
restore_from_version_no,
|
||||
creator,
|
||||
created_at
|
||||
)
|
||||
SELECT #{snapshot.id},
|
||||
#{snapshot.agentId},
|
||||
#{snapshot.userId},
|
||||
COALESCE(MAX(version_no), 0) + 1,
|
||||
#{snapshot.snapshotData},
|
||||
#{snapshot.changedFields},
|
||||
#{snapshot.source},
|
||||
#{snapshot.restoreFromSnapshotId},
|
||||
#{snapshot.restoreFromVersionNo},
|
||||
#{snapshot.creator},
|
||||
#{snapshot.createdAt}
|
||||
FROM ai_agent_snapshot
|
||||
WHERE agent_id = #{snapshot.agentId}
|
||||
</insert>
|
||||
|
||||
<delete id="deleteOlderThanKeepLimit">
|
||||
DELETE FROM ai_agent_snapshot
|
||||
WHERE agent_id = #{agentId}
|
||||
AND id NOT IN (
|
||||
SELECT id
|
||||
FROM (
|
||||
SELECT id
|
||||
FROM ai_agent_snapshot
|
||||
WHERE agent_id = #{agentId}
|
||||
ORDER BY version_no DESC
|
||||
LIMIT #{keepLimit}
|
||||
) retained_snapshots
|
||||
)
|
||||
</delete>
|
||||
|
||||
</mapper>
|
||||
@@ -0,0 +1,44 @@
|
||||
package xiaozhi.modules.agent.dto;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.assertEquals;
|
||||
import static org.junit.jupiter.api.Assertions.assertTrue;
|
||||
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.Map;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
class AgentUpdateDTOTest {
|
||||
|
||||
@Test
|
||||
void functionInfoAcceptsJsonStringParamInfo() {
|
||||
AgentUpdateDTO.FunctionInfo info = new AgentUpdateDTO.FunctionInfo();
|
||||
|
||||
info.setParamInfo("{\"api_key\":\"abc\",\"max_results\":5}");
|
||||
|
||||
assertEquals("abc", info.getParamInfo().get("api_key"));
|
||||
assertEquals(5, info.getParamInfo().get("max_results"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void functionInfoNormalizesMapKeysToStrings() {
|
||||
AgentUpdateDTO.FunctionInfo info = new AgentUpdateDTO.FunctionInfo();
|
||||
Map<Object, Object> params = new LinkedHashMap<>();
|
||||
params.put("api_key", "abc");
|
||||
params.put(42, true);
|
||||
|
||||
info.setParamInfo(params);
|
||||
|
||||
assertEquals("abc", info.getParamInfo().get("api_key"));
|
||||
assertEquals(true, info.getParamInfo().get("42"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void functionInfoFallsBackToEmptyMapForBlankParamInfo() {
|
||||
AgentUpdateDTO.FunctionInfo info = new AgentUpdateDTO.FunctionInfo();
|
||||
|
||||
info.setParamInfo(" ");
|
||||
|
||||
assertTrue(info.getParamInfo().isEmpty());
|
||||
}
|
||||
}
|
||||
+450
@@ -0,0 +1,450 @@
|
||||
package xiaozhi.modules.agent.service.impl;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.assertEquals;
|
||||
import static org.junit.jupiter.api.Assertions.assertFalse;
|
||||
import static org.junit.jupiter.api.Assertions.assertNotNull;
|
||||
import static org.junit.jupiter.api.Assertions.assertNotEquals;
|
||||
import static org.junit.jupiter.api.Assertions.assertTrue;
|
||||
import static org.mockito.ArgumentMatchers.any;
|
||||
import static org.mockito.ArgumentMatchers.anyInt;
|
||||
import static org.mockito.ArgumentMatchers.argThat;
|
||||
import static org.mockito.Mockito.inOrder;
|
||||
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 com.baomidou.mybatisplus.extension.plugins.pagination.Page;
|
||||
import java.lang.reflect.Method;
|
||||
import java.nio.file.Files;
|
||||
import java.nio.file.Path;
|
||||
import java.util.Date;
|
||||
import java.util.HashMap;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.mockito.InOrder;
|
||||
import org.springframework.test.util.ReflectionTestUtils;
|
||||
import org.springframework.transaction.annotation.Transactional;
|
||||
|
||||
import xiaozhi.common.utils.JsonUtils;
|
||||
import xiaozhi.modules.agent.Enums.AgentSnapshotField;
|
||||
import xiaozhi.modules.agent.dao.AgentDao;
|
||||
import xiaozhi.modules.agent.dao.AgentSnapshotDao;
|
||||
import xiaozhi.modules.agent.dao.AgentTagDao;
|
||||
import xiaozhi.modules.agent.dto.AgentSnapshotDataDTO;
|
||||
import xiaozhi.modules.agent.dto.AgentSnapshotPageDTO;
|
||||
import xiaozhi.modules.agent.dto.AgentSnapshotTagDTO;
|
||||
import xiaozhi.modules.agent.dto.AgentUpdateDTO;
|
||||
import xiaozhi.modules.agent.dto.ContextProviderDTO;
|
||||
import xiaozhi.modules.agent.entity.AgentEntity;
|
||||
import xiaozhi.modules.agent.entity.AgentSnapshotEntity;
|
||||
import xiaozhi.modules.agent.entity.AgentTagEntity;
|
||||
import xiaozhi.modules.agent.service.AgentContextProviderService;
|
||||
import xiaozhi.modules.agent.service.AgentSnapshotService;
|
||||
import xiaozhi.modules.agent.vo.AgentSnapshotVO;
|
||||
import xiaozhi.modules.agent.vo.AgentInfoVO;
|
||||
import xiaozhi.modules.correctword.service.CorrectWordFileService;
|
||||
|
||||
class AgentSnapshotServiceImplTest {
|
||||
|
||||
@Test
|
||||
@SuppressWarnings("unchecked")
|
||||
void normalizeSortedJsonListTreatsDtoAndMapShapesEqually() throws Exception {
|
||||
AgentSnapshotServiceImpl service = new AgentSnapshotServiceImpl(null, null, null, null, null, null, null, null,
|
||||
null, null);
|
||||
Method method = AgentSnapshotServiceImpl.class.getDeclaredMethod("normalizeSortedJsonList", List.class);
|
||||
method.setAccessible(true);
|
||||
|
||||
ContextProviderDTO dto = new ContextProviderDTO();
|
||||
dto.setUrl("https://example.com/context");
|
||||
dto.setHeaders(Map.of("Authorization", "Bearer token"));
|
||||
|
||||
Map<String, Object> map = new LinkedHashMap<>();
|
||||
map.put("headers", Map.of("Authorization", "Bearer token"));
|
||||
map.put("url", "https://example.com/context");
|
||||
|
||||
assertEquals(method.invoke(service, List.of(dto)), method.invoke(service, List.of(map)));
|
||||
}
|
||||
|
||||
@Test
|
||||
void redactSnapshotDataMasksSensitiveFunctionAndHeaderValues() throws Exception {
|
||||
AgentSnapshotServiceImpl service = new AgentSnapshotServiceImpl(null, null, null, null, null, null, null, null,
|
||||
null, null);
|
||||
Method method = AgentSnapshotServiceImpl.class.getDeclaredMethod("redactSnapshotData",
|
||||
AgentSnapshotDataDTO.class);
|
||||
method.setAccessible(true);
|
||||
|
||||
AgentSnapshotDataDTO data = new AgentSnapshotDataDTO();
|
||||
AgentUpdateDTO.FunctionInfo function = new AgentUpdateDTO.FunctionInfo();
|
||||
function.setPluginId("rag");
|
||||
function.setParamInfo(Map.of("api_key", "secret-key", "max_tokens", 32));
|
||||
ContextProviderDTO provider = new ContextProviderDTO();
|
||||
provider.setUrl("https://example.com/context");
|
||||
provider.setHeaders(Map.of("Authorization", "Bearer secret-token", "Accept", "application/json"));
|
||||
data.setFunctions(List.of(function));
|
||||
data.setContextProviders(List.of(provider));
|
||||
|
||||
AgentSnapshotDataDTO redacted = (AgentSnapshotDataDTO) method.invoke(service, data);
|
||||
String json = JsonUtils.toJsonString(redacted);
|
||||
|
||||
assertFalse(json.contains("secret-key"));
|
||||
assertFalse(json.contains("secret-token"));
|
||||
assertTrue(json.contains("__SNAPSHOT_SECRET_REDACTED__"));
|
||||
assertTrue(json.contains("max_tokens"));
|
||||
assertTrue(json.contains("application/json"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void preserveCurrentSensitiveValuesKeepsCurrentSecretsWhenRestoringSnapshot() throws Exception {
|
||||
AgentSnapshotServiceImpl service = new AgentSnapshotServiceImpl(null, null, null, null, null, null, null, null,
|
||||
null, null);
|
||||
Method method = AgentSnapshotServiceImpl.class.getDeclaredMethod("preserveCurrentSensitiveValues",
|
||||
AgentSnapshotDataDTO.class, AgentSnapshotDataDTO.class);
|
||||
method.setAccessible(true);
|
||||
|
||||
AgentSnapshotDataDTO snapshot = buildSnapshot("old-key", "old-token", "old-value");
|
||||
AgentSnapshotDataDTO current = buildSnapshot("current-key", "current-token", "current-value");
|
||||
|
||||
AgentSnapshotDataDTO restored = (AgentSnapshotDataDTO) method.invoke(service, snapshot, current);
|
||||
|
||||
assertEquals("current-key", restored.getFunctions().get(0).getParamInfo().get("api_key"));
|
||||
assertEquals("old-value", restored.getFunctions().get(0).getParamInfo().get("description"));
|
||||
assertEquals("current-token", restored.getContextProviders().get(0).getHeaders().get("Authorization"));
|
||||
}
|
||||
|
||||
@Test
|
||||
@SuppressWarnings("unchecked")
|
||||
void getChangedFieldsFollowsCentralFieldOrder() throws Exception {
|
||||
AgentSnapshotServiceImpl service = new AgentSnapshotServiceImpl(null, null, null, null, null, null, null, null,
|
||||
null, null);
|
||||
Method method = AgentSnapshotServiceImpl.class.getDeclaredMethod("getChangedFields",
|
||||
AgentSnapshotDataDTO.class, AgentUpdateDTO.class);
|
||||
method.setAccessible(true);
|
||||
|
||||
AgentSnapshotDataDTO current = new AgentSnapshotDataDTO();
|
||||
current.setAgentName("old-name");
|
||||
current.setCorrectWordFileIds(List.of("b", "a"));
|
||||
|
||||
AgentUpdateDTO pendingUpdate = new AgentUpdateDTO();
|
||||
pendingUpdate.setAgentName("new-name");
|
||||
pendingUpdate.setCorrectWordFileIds(List.of("a", "b"));
|
||||
AgentUpdateDTO.FunctionInfo function = new AgentUpdateDTO.FunctionInfo();
|
||||
function.setPluginId("rag");
|
||||
function.setParamInfo(Map.of("api_key", "secret-key"));
|
||||
pendingUpdate.setFunctions(List.of(function));
|
||||
|
||||
List<String> changedFields = (List<String>) method.invoke(service, current, pendingUpdate);
|
||||
|
||||
assertEquals(List.of(AgentSnapshotField.AGENT_NAME.getFieldName(),
|
||||
AgentSnapshotField.FUNCTIONS.getFieldName()), changedFields);
|
||||
}
|
||||
|
||||
@Test
|
||||
@SuppressWarnings("unchecked")
|
||||
void tagIdsAreNotExposedAsIndependentSnapshotFields() throws Exception {
|
||||
AgentSnapshotServiceImpl service = new AgentSnapshotServiceImpl(null, null, null, null, null, null, null, null,
|
||||
null, null);
|
||||
Method method = AgentSnapshotServiceImpl.class.getDeclaredMethod("getChangedFields",
|
||||
AgentSnapshotDataDTO.class, AgentUpdateDTO.class);
|
||||
method.setAccessible(true);
|
||||
|
||||
AgentSnapshotDataDTO current = new AgentSnapshotDataDTO();
|
||||
current.setTagNames(List.of("alpha", "beta"));
|
||||
|
||||
AgentUpdateDTO pendingUpdate = new AgentUpdateDTO();
|
||||
pendingUpdate.setTagNames(List.of("beta", "alpha"));
|
||||
pendingUpdate.setTagIds(List.of("new-id"));
|
||||
|
||||
List<String> changedFields = (List<String>) method.invoke(service, current, pendingUpdate);
|
||||
|
||||
assertFalse(AgentSnapshotField.names().contains("tagIds"));
|
||||
assertFalse(changedFields.contains("tagIds"));
|
||||
assertTrue(changedFields.isEmpty());
|
||||
|
||||
pendingUpdate.setTagNames(List.of("alpha", "gamma"));
|
||||
changedFields = (List<String>) method.invoke(service, current, pendingUpdate);
|
||||
|
||||
assertEquals(List.of(AgentSnapshotField.TAG_NAMES.getFieldName()), changedFields);
|
||||
}
|
||||
|
||||
@Test
|
||||
void tagOnlySavePathCreatesSnapshot() throws Exception {
|
||||
String controller = Files.readString(Path.of("src/main/java/xiaozhi/modules/agent/controller/AgentController.java"));
|
||||
String roleConfig = Files.readString(Path.of("../manager-web/src/views/roleConfig.vue"));
|
||||
|
||||
assertTrue(controller.contains("agentService.updateAgentById(id, dto);"));
|
||||
assertFalse(controller.contains("agentService.updateAgentById(id, dto, false);"));
|
||||
assertTrue(roleConfig.contains("configData.tagNames = tagNames;"));
|
||||
assertFalse(roleConfig.contains("this.handleSaveAgentTags(agentId, tagNames)"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void firstAgentSaveKeepsInitialBaselineBeforePersistingNewState() {
|
||||
AgentDao agentDao = mock(AgentDao.class);
|
||||
AgentContextProviderService contextProviderService = mock(AgentContextProviderService.class);
|
||||
CorrectWordFileService correctWordFileService = mock(CorrectWordFileService.class);
|
||||
AgentSnapshotService snapshotService = mock(AgentSnapshotService.class);
|
||||
AgentServiceImpl service = new AgentServiceImpl(agentDao, null, null, null, null, null, null, null,
|
||||
null, null, contextProviderService, null, correctWordFileService, snapshotService);
|
||||
ReflectionTestUtils.setField(service, "baseDao", agentDao);
|
||||
|
||||
String agentId = "agent-id";
|
||||
AgentEntity lockedAgent = new AgentEntity();
|
||||
lockedAgent.setId(agentId);
|
||||
AgentInfoVO currentAgent = new AgentInfoVO();
|
||||
currentAgent.setId(agentId);
|
||||
currentAgent.setAgentName("old-name");
|
||||
AgentUpdateDTO update = new AgentUpdateDTO();
|
||||
update.setAgentName("new-name");
|
||||
|
||||
when(agentDao.selectByIdForUpdate(agentId)).thenReturn(lockedAgent);
|
||||
when(agentDao.selectAgentInfoById(agentId)).thenReturn(currentAgent);
|
||||
when(snapshotService.getCurrentVersionNo(agentId)).thenReturn(0);
|
||||
when(contextProviderService.getByAgentId(agentId)).thenReturn(null);
|
||||
when(correctWordFileService.getAgentCorrectWordFileIds(agentId)).thenReturn(List.of());
|
||||
when(agentDao.updateById(any())).thenReturn(1);
|
||||
|
||||
service.updateAgentById(agentId, update);
|
||||
|
||||
InOrder inOrder = inOrder(agentDao, snapshotService);
|
||||
inOrder.verify(snapshotService).createSnapshot(agentId, "initial");
|
||||
inOrder.verify(agentDao).updateById(argThat(agent -> "new-name".equals(agent.getAgentName())));
|
||||
inOrder.verify(snapshotService).createSnapshot(agentId, "config");
|
||||
}
|
||||
|
||||
@Test
|
||||
void agentSnapshotFieldAppliesRestorableAgentFields() {
|
||||
AgentSnapshotDataDTO data = new AgentSnapshotDataDTO();
|
||||
data.setAgentName("restored-name");
|
||||
data.setLlmModelId("llm-new");
|
||||
data.setSystemPrompt("restored prompt");
|
||||
data.setSort(7);
|
||||
|
||||
AgentEntity agent = new AgentEntity();
|
||||
for (AgentSnapshotField field : AgentSnapshotField.values()) {
|
||||
field.applyTo(agent, data);
|
||||
}
|
||||
|
||||
assertEquals("restored-name", agent.getAgentName());
|
||||
assertEquals("llm-new", agent.getLlmModelId());
|
||||
assertEquals("restored prompt", agent.getSystemPrompt());
|
||||
assertEquals(7, agent.getSort());
|
||||
}
|
||||
|
||||
@Test
|
||||
void restoreSnapshotRollsBackForAnyException() throws Exception {
|
||||
Method method = AgentSnapshotServiceImpl.class.getMethod("restoreSnapshot", String.class, String.class);
|
||||
|
||||
Transactional transactional = method.getAnnotation(Transactional.class);
|
||||
|
||||
assertNotNull(transactional);
|
||||
assertTrue(List.of(transactional.rollbackFor()).contains(Exception.class));
|
||||
}
|
||||
|
||||
@Test
|
||||
void deleteSnapshotRollsBackForAnyException() throws Exception {
|
||||
Method method = AgentSnapshotServiceImpl.class.getMethod("deleteSnapshot", String.class, String.class);
|
||||
|
||||
Transactional transactional = method.getAnnotation(Transactional.class);
|
||||
|
||||
assertNotNull(transactional);
|
||||
assertTrue(List.of(transactional.rollbackFor()).contains(Exception.class));
|
||||
}
|
||||
|
||||
@Test
|
||||
void snapshotInsertAllocatesVersionInSingleSqlStatement() throws Exception {
|
||||
String xml = normalizeWhitespace(Files.readString(Path.of("src/main/resources/mapper/agent/AgentSnapshotDao.xml")));
|
||||
String dao = Files.readString(Path.of("src/main/java/xiaozhi/modules/agent/dao/AgentSnapshotDao.java"));
|
||||
|
||||
assertTrue(dao.contains("insertWithNextVersion"));
|
||||
assertTrue(xml.contains("<insert id=\"insertWithNextVersion\">"));
|
||||
assertTrue(xml.contains("INSERT INTO ai_agent_snapshot"));
|
||||
assertTrue(xml.contains("COALESCE(MAX(version_no), 0) + 1"));
|
||||
assertFalse(xml.contains("LAST_INSERT_ID"));
|
||||
assertFalse(xml.contains("ai_agent_snapshot_sequence"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void snapshotMigrationIsMergedIntoSingleChangeLog() throws Exception {
|
||||
String sql = Files.readString(Path.of("src/main/resources/db/changelog/202607071530.sql"));
|
||||
String master = Files.readString(Path.of("src/main/resources/db/changelog/db.changelog-master.yaml"));
|
||||
|
||||
assertTrue(sql.contains("restore_from_snapshot_id"));
|
||||
assertTrue(sql.contains("restore_from_version_no"));
|
||||
assertTrue(sql.contains("idx_snapshot_user_created_at"));
|
||||
assertTrue(sql.contains("DEFAULT (JSON_OBJECT())"));
|
||||
assertTrue(master.contains("202607071530.sql"));
|
||||
assertFalse(master.contains("202607081150.sql"));
|
||||
assertFalse(master.contains("202607081230.sql"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void selectNextSnapshotLoadsRestoreTraceColumns() throws Exception {
|
||||
String xml = normalizeWhitespace(Files.readString(Path.of("src/main/resources/mapper/agent/AgentSnapshotDao.xml")));
|
||||
|
||||
assertTrue(xml.contains("restore_from_snapshot_id AS restoreFromSnapshotId"));
|
||||
assertTrue(xml.contains("restore_from_version_no AS restoreFromVersionNo"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void retentionPruneSqlKeepsLatestVersions() throws Exception {
|
||||
String xml = Files.readString(Path.of("src/main/resources/mapper/agent/AgentSnapshotDao.xml"));
|
||||
|
||||
assertTrue(xml.contains("deleteOlderThanKeepLimit"));
|
||||
assertTrue(xml.contains("ORDER BY version_no DESC"));
|
||||
assertTrue(xml.contains("LIMIT #{keepLimit}"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void agentInfoMapperLoadsFunctionsWithJoinInsteadOfNestedSelect() throws Exception {
|
||||
String xml = Files.readString(Path.of("src/main/resources/mapper/agent/AgentDao.xml"));
|
||||
|
||||
assertTrue(xml.contains("LEFT JOIN ai_agent_plugin_mapping f ON f.agent_id = a.id"));
|
||||
assertFalse(xml.contains("selectAgentFunctionsByAgentId"));
|
||||
assertFalse(xml.contains("select=\"selectAgentFunctions"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void toVOIncludesRestoreTraceFields() throws Exception {
|
||||
AgentSnapshotServiceImpl service = new AgentSnapshotServiceImpl(null, null, null, null, null, null, null, null,
|
||||
null, null);
|
||||
Method method = AgentSnapshotServiceImpl.class.getDeclaredMethod("toVO",
|
||||
AgentSnapshotEntity.class, boolean.class);
|
||||
method.setAccessible(true);
|
||||
|
||||
AgentSnapshotEntity entity = new AgentSnapshotEntity();
|
||||
entity.setId("snapshot-id");
|
||||
entity.setAgentId("agent-id");
|
||||
entity.setVersionNo(3);
|
||||
entity.setChangedFields("[]");
|
||||
entity.setSource("restore");
|
||||
entity.setRestoreFromSnapshotId("target-id");
|
||||
entity.setRestoreFromVersionNo(1);
|
||||
|
||||
AgentSnapshotVO vo = (AgentSnapshotVO) method.invoke(service, entity, false);
|
||||
|
||||
assertEquals("target-id", vo.getRestoreFromSnapshotId());
|
||||
assertEquals(1, vo.getRestoreFromVersionNo());
|
||||
}
|
||||
|
||||
@Test
|
||||
void pageDoesNotCreateInitialSnapshotWhenAgentHasNoHistory() {
|
||||
AgentSnapshotDao snapshotDao = mock(AgentSnapshotDao.class);
|
||||
AgentDao agentDao = mock(AgentDao.class);
|
||||
AgentSnapshotServiceImpl service = new AgentSnapshotServiceImpl(snapshotDao, agentDao, null, null, null,
|
||||
null, null, null, null, null);
|
||||
|
||||
when(snapshotDao.selectPage(any(), any())).thenReturn(new Page<AgentSnapshotEntity>(1, 10));
|
||||
|
||||
service.page("agent-id", new AgentSnapshotPageDTO());
|
||||
|
||||
verify(agentDao, never()).selectByIdForUpdate(any());
|
||||
verify(snapshotDao, never()).selectMaxVersionNo(any());
|
||||
verify(snapshotDao, never()).insertWithNextVersion(any());
|
||||
verify(snapshotDao, never()).deleteOlderThanKeepLimit(any(), anyInt());
|
||||
}
|
||||
|
||||
@Test
|
||||
void getCurrentVersionNoReturnsPersistedMaxVersion() {
|
||||
AgentSnapshotDao snapshotDao = mock(AgentSnapshotDao.class);
|
||||
AgentSnapshotServiceImpl service = new AgentSnapshotServiceImpl(snapshotDao, null, null, null, null,
|
||||
null, null, null, null, null);
|
||||
|
||||
when(snapshotDao.selectMaxVersionNo("agent-id")).thenReturn(7);
|
||||
|
||||
assertEquals(7, service.getCurrentVersionNo("agent-id"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void createSnapshotStoresCurrentStateAsInitialVersionWhenHistoryIsEmpty() {
|
||||
AgentSnapshotDao snapshotDao = mock(AgentSnapshotDao.class);
|
||||
AgentDao agentDao = mock(AgentDao.class);
|
||||
AgentTagDao tagDao = mock(AgentTagDao.class);
|
||||
AgentContextProviderService contextProviderService = mock(AgentContextProviderService.class);
|
||||
CorrectWordFileService correctWordFileService = mock(CorrectWordFileService.class);
|
||||
AgentSnapshotServiceImpl service = new AgentSnapshotServiceImpl(snapshotDao, agentDao, tagDao, null, null,
|
||||
contextProviderService, null, null, null, correctWordFileService);
|
||||
|
||||
AgentEntity lockedAgent = new AgentEntity();
|
||||
lockedAgent.setId("agent-id");
|
||||
AgentInfoVO agentInfo = new AgentInfoVO();
|
||||
agentInfo.setId("agent-id");
|
||||
agentInfo.setUserId(7L);
|
||||
agentInfo.setAgentName("agent");
|
||||
when(agentDao.selectByIdForUpdate("agent-id")).thenReturn(lockedAgent);
|
||||
when(agentDao.selectAgentInfoById("agent-id")).thenReturn(agentInfo);
|
||||
when(snapshotDao.selectLatestSnapshot("agent-id")).thenReturn(null);
|
||||
when(snapshotDao.insertWithNextVersion(any())).thenReturn(1);
|
||||
when(correctWordFileService.getAgentCorrectWordFileIds("agent-id")).thenReturn(List.of());
|
||||
when(contextProviderService.getByAgentId("agent-id")).thenReturn(null);
|
||||
when(tagDao.selectByAgentId("agent-id")).thenReturn(List.of());
|
||||
|
||||
service.createSnapshot("agent-id", "config");
|
||||
|
||||
verify(snapshotDao).insertWithNextVersion(argThat(snapshot -> "config".equals(snapshot.getSource())
|
||||
&& snapshot.getVersionNo() == null
|
||||
&& snapshot.getChangedFields().contains("initial")));
|
||||
verify(snapshotDao).deleteOlderThanKeepLimit("agent-id", 100);
|
||||
}
|
||||
|
||||
@Test
|
||||
void restoreTagDoesNotReviveSoftDeletedSnapshotTagId() throws Exception {
|
||||
AgentTagDao tagDao = mock(AgentTagDao.class);
|
||||
AgentSnapshotServiceImpl service = new AgentSnapshotServiceImpl(null, null, tagDao, null, null, null, null, null,
|
||||
null, null);
|
||||
Method method = AgentSnapshotServiceImpl.class.getDeclaredMethod("restoreTag",
|
||||
AgentSnapshotTagDTO.class, Date.class);
|
||||
method.setAccessible(true);
|
||||
|
||||
AgentTagEntity deletedTag = new AgentTagEntity();
|
||||
deletedTag.setId("deleted-tag-id");
|
||||
deletedTag.setTagName("archived");
|
||||
deletedTag.setDeleted(1);
|
||||
when(tagDao.selectById("deleted-tag-id")).thenReturn(deletedTag);
|
||||
when(tagDao.selectOne(any())).thenReturn(null);
|
||||
AtomicReference<AgentTagEntity> insertedTag = new AtomicReference<>();
|
||||
when(tagDao.insert(any())).thenAnswer(invocation -> {
|
||||
insertedTag.set(invocation.getArgument(0));
|
||||
return 1;
|
||||
});
|
||||
|
||||
AgentSnapshotTagDTO snapshotTag = new AgentSnapshotTagDTO();
|
||||
snapshotTag.setId("deleted-tag-id");
|
||||
snapshotTag.setTagName("archived");
|
||||
|
||||
String restoredTagId = (String) method.invoke(service, snapshotTag, new Date());
|
||||
|
||||
verify(tagDao, never()).updateById(any());
|
||||
assertNotEquals("deleted-tag-id", restoredTagId);
|
||||
assertNotNull(insertedTag.get());
|
||||
assertEquals(restoredTagId, insertedTag.get().getId());
|
||||
assertEquals(0, insertedTag.get().getDeleted());
|
||||
}
|
||||
|
||||
private AgentSnapshotDataDTO buildSnapshot(String apiKey, String authToken, String description) {
|
||||
AgentSnapshotDataDTO data = new AgentSnapshotDataDTO();
|
||||
|
||||
AgentUpdateDTO.FunctionInfo function = new AgentUpdateDTO.FunctionInfo();
|
||||
function.setPluginId("rag");
|
||||
HashMap<String, Object> params = new HashMap<>();
|
||||
params.put("api_key", apiKey);
|
||||
params.put("description", description);
|
||||
function.setParamInfo(params);
|
||||
|
||||
ContextProviderDTO provider = new ContextProviderDTO();
|
||||
provider.setUrl("https://example.com/context");
|
||||
provider.setHeaders(Map.of("Authorization", authToken));
|
||||
|
||||
data.setFunctions(List.of(function));
|
||||
data.setContextProviders(List.of(provider));
|
||||
return data;
|
||||
}
|
||||
|
||||
private String normalizeWhitespace(String value) {
|
||||
return value.replaceAll("\\s+", " ").trim();
|
||||
}
|
||||
}
|
||||
+42
@@ -0,0 +1,42 @@
|
||||
package xiaozhi.modules.agent.typehandler;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.assertEquals;
|
||||
import static org.junit.jupiter.api.Assertions.assertInstanceOf;
|
||||
import static org.junit.jupiter.api.Assertions.assertTrue;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import org.apache.ibatis.type.TypeHandler;
|
||||
import org.apache.ibatis.type.TypeHandlerRegistry;
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
import xiaozhi.modules.agent.dto.ContextProviderDTO;
|
||||
|
||||
class ContextProviderListTypeHandlerTest {
|
||||
|
||||
private final ContextProviderListTypeHandler handler = new ContextProviderListTypeHandler();
|
||||
|
||||
@Test
|
||||
void parseKeepsContextProviderDtoElementType() {
|
||||
List<ContextProviderDTO> providers = handler
|
||||
.parse("[{\"url\":\"https://example.com/context\",\"headers\":{\"Authorization\":\"Bearer token\"}}]");
|
||||
|
||||
assertEquals(1, providers.size());
|
||||
assertInstanceOf(ContextProviderDTO.class, providers.get(0));
|
||||
assertEquals("https://example.com/context", providers.get(0).getUrl());
|
||||
assertEquals("Bearer token", providers.get(0).getHeaders().get("Authorization"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void parseBlankJsonAsEmptyList() {
|
||||
assertTrue(handler.parse(" ").isEmpty());
|
||||
}
|
||||
|
||||
@Test
|
||||
void myBatisCanInstantiateHandlerForListField() {
|
||||
TypeHandler<?> typeHandler = new TypeHandlerRegistry().getInstance(List.class,
|
||||
ContextProviderListTypeHandler.class);
|
||||
|
||||
assertInstanceOf(ContextProviderListTypeHandler.class, typeHandler);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user