feat: AI 助手模型管理(CRUD、会话绑定、自动化验收)
- Flyway V6:模型表、默认种子、ai_chat_session.model_id 回填 - 管理端 /api/admin/ai/models;用户端选模与新会话绑定 - 禁用模型后会话只读(5009);chat 按会话 modelCode 调 DashScope - 前端 AiModels 页、助手模型选择器与 localStorage - 单测、AiModelManagementApiIT、Playwright E2E;e2e Profile 关闭注解限流 - 技术实施文档与验收报告;移除过期冲突分析文档 Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
+54
@@ -0,0 +1,54 @@
|
||||
package fun.nojava.module.blog.controller;
|
||||
|
||||
import fun.nojava.common.model.ApiResult;
|
||||
import fun.nojava.module.blog.dto.AiAssistantModelSaveDTO;
|
||||
import fun.nojava.module.blog.dto.AiAssistantModelVO;
|
||||
import fun.nojava.module.blog.service.AiAssistantModelService;
|
||||
import io.swagger.v3.oas.annotations.Operation;
|
||||
import io.swagger.v3.oas.annotations.tags.Tag;
|
||||
import jakarta.validation.Valid;
|
||||
import lombok.RequiredArgsConstructor;
|
||||
import org.springframework.web.bind.annotation.*;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
@Tag(name = "管理员-AI 助手模型")
|
||||
@RestController
|
||||
@RequestMapping("/api/admin/ai/models")
|
||||
@RequiredArgsConstructor
|
||||
public class AdminAiModelController {
|
||||
|
||||
private final AiAssistantModelService aiAssistantModelService;
|
||||
|
||||
@Operation(summary = "模型列表")
|
||||
@GetMapping
|
||||
public ApiResult<List<AiAssistantModelVO>> list() {
|
||||
return ApiResult.success(aiAssistantModelService.listAllForAdmin());
|
||||
}
|
||||
|
||||
@Operation(summary = "模型详情")
|
||||
@GetMapping("/{id}")
|
||||
public ApiResult<AiAssistantModelVO> get(@PathVariable Long id) {
|
||||
return ApiResult.success(aiAssistantModelService.getByIdForAdmin(id));
|
||||
}
|
||||
|
||||
@Operation(summary = "新增模型")
|
||||
@PostMapping
|
||||
public ApiResult<AiAssistantModelVO> create(@Valid @RequestBody AiAssistantModelSaveDTO dto) {
|
||||
return ApiResult.success(aiAssistantModelService.create(dto));
|
||||
}
|
||||
|
||||
@Operation(summary = "更新模型")
|
||||
@PutMapping("/{id}")
|
||||
public ApiResult<Void> update(@PathVariable Long id, @Valid @RequestBody AiAssistantModelSaveDTO dto) {
|
||||
aiAssistantModelService.update(id, dto);
|
||||
return ApiResult.success();
|
||||
}
|
||||
|
||||
@Operation(summary = "逻辑删除模型")
|
||||
@DeleteMapping("/{id}")
|
||||
public ApiResult<Void> delete(@PathVariable Long id) {
|
||||
aiAssistantModelService.deleteLogical(id);
|
||||
return ApiResult.success();
|
||||
}
|
||||
}
|
||||
+12
-3
@@ -21,10 +21,19 @@ public class AiChatController {
|
||||
|
||||
private final AiChatService aiChatService;
|
||||
|
||||
@Operation(summary = "可选 AI 助手模型列表")
|
||||
@GetMapping("/models")
|
||||
public ApiResult<List<AiAssistantModelAppVO>> listModels(@AuthenticationPrincipal Long userId) {
|
||||
return ApiResult.success(aiChatService.listModelsForApp());
|
||||
}
|
||||
|
||||
@Operation(summary = "创建 AI 对话会话")
|
||||
@PostMapping("/sessions")
|
||||
public ApiResult<CreateAiChatSessionVO> createSession(@AuthenticationPrincipal Long userId) {
|
||||
return ApiResult.success(aiChatService.createSession(userId));
|
||||
public ApiResult<CreateAiChatSessionVO> createSession(
|
||||
@AuthenticationPrincipal Long userId,
|
||||
@RequestBody(required = false) CreateAiChatSessionDTO body) {
|
||||
Long modelId = body == null ? null : body.getModelId();
|
||||
return ApiResult.success(aiChatService.createSession(userId, modelId));
|
||||
}
|
||||
|
||||
@Operation(summary = "分页查询当前用户的 AI 会话列表")
|
||||
@@ -38,7 +47,7 @@ public class AiChatController {
|
||||
|
||||
@Operation(summary = "查询会话全部消息")
|
||||
@GetMapping("/sessions/{sessionId}/messages")
|
||||
public ApiResult<List<AiChatMessageVO>> listMessages(
|
||||
public ApiResult<AiChatSessionMessagesVO> listMessages(
|
||||
@AuthenticationPrincipal Long userId,
|
||||
@PathVariable Long sessionId) {
|
||||
return ApiResult.success(aiChatService.listMessages(userId, sessionId));
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
package fun.nojava.module.blog.dto;
|
||||
|
||||
import lombok.Builder;
|
||||
import lombok.Data;
|
||||
|
||||
@Data
|
||||
@Builder
|
||||
public class AiAssistantModelAppVO {
|
||||
|
||||
private Long id;
|
||||
private String displayName;
|
||||
private String description;
|
||||
private Boolean isDefault;
|
||||
private Boolean allowImageUpload;
|
||||
private Boolean allowVoiceUpload;
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
package fun.nojava.module.blog.dto;
|
||||
|
||||
import jakarta.validation.constraints.NotBlank;
|
||||
import jakarta.validation.constraints.Size;
|
||||
import lombok.Data;
|
||||
|
||||
@Data
|
||||
public class AiAssistantModelSaveDTO {
|
||||
|
||||
@NotBlank(message = "模型编码不能为空")
|
||||
@Size(min = 1, max = 64, message = "模型编码长度须在 1~64 字符之间")
|
||||
private String modelCode;
|
||||
|
||||
@NotBlank(message = "展示名称不能为空")
|
||||
@Size(min = 1, max = 32, message = "展示名称长度须在 1~32 字符之间")
|
||||
private String displayName;
|
||||
|
||||
@Size(max = 200, message = "描述最长 200 字符")
|
||||
private String description;
|
||||
|
||||
private Boolean enabled = true;
|
||||
|
||||
private Boolean allowImageUpload = false;
|
||||
|
||||
private Boolean allowVoiceUpload = false;
|
||||
|
||||
private Boolean isDefault = false;
|
||||
|
||||
private Integer sort = 0;
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
package fun.nojava.module.blog.dto;
|
||||
|
||||
import lombok.Builder;
|
||||
import lombok.Data;
|
||||
|
||||
import java.time.LocalDateTime;
|
||||
|
||||
@Data
|
||||
@Builder
|
||||
public class AiAssistantModelVO {
|
||||
|
||||
private Long id;
|
||||
private String modelCode;
|
||||
private String displayName;
|
||||
private String description;
|
||||
private Boolean enabled;
|
||||
private Boolean allowImageUpload;
|
||||
private Boolean allowVoiceUpload;
|
||||
private Boolean isDefault;
|
||||
private Integer sort;
|
||||
private LocalDateTime createdAt;
|
||||
private LocalDateTime updatedAt;
|
||||
}
|
||||
@@ -17,6 +17,8 @@ public class AiChatResultVO {
|
||||
|
||||
private String model;
|
||||
|
||||
private String modelDisplayName;
|
||||
|
||||
private Integer promptTokens;
|
||||
|
||||
private Integer completionTokens;
|
||||
|
||||
@@ -14,4 +14,7 @@ public class AiChatSendDTO {
|
||||
@NotBlank(message = "问题不能为空")
|
||||
@Size(max = 2000, message = "问题长度不能超过 2000 字符")
|
||||
private String message;
|
||||
|
||||
/** 可选;与会话已绑定模型不一致时忽略 */
|
||||
private Long modelId;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
package fun.nojava.module.blog.dto;
|
||||
|
||||
import fun.nojava.module.blog.enums.AiChatReadOnlyReason;
|
||||
import lombok.Builder;
|
||||
import lombok.Data;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
@Data
|
||||
@Builder
|
||||
public class AiChatSessionMessagesVO {
|
||||
|
||||
private List<AiChatMessageVO> messages;
|
||||
private Boolean modelEnabled;
|
||||
private AiChatReadOnlyReason readOnlyReason;
|
||||
private Integer roundCount;
|
||||
private Integer maxRounds;
|
||||
}
|
||||
@@ -16,4 +16,10 @@ public class AiChatSessionVO {
|
||||
private Integer roundCount;
|
||||
|
||||
private LocalDateTime lastMessageAt;
|
||||
|
||||
private Long modelId;
|
||||
|
||||
private String modelDisplayName;
|
||||
|
||||
private Boolean modelEnabled;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
package fun.nojava.module.blog.dto;
|
||||
|
||||
import lombok.Data;
|
||||
|
||||
@Data
|
||||
public class CreateAiChatSessionDTO {
|
||||
|
||||
private Long modelId;
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package fun.nojava.module.blog.entity;
|
||||
|
||||
import com.baomidou.mybatisplus.annotation.*;
|
||||
import lombok.Data;
|
||||
|
||||
import java.time.LocalDateTime;
|
||||
|
||||
@Data
|
||||
@TableName("ai_assistant_model")
|
||||
public class AiAssistantModel {
|
||||
|
||||
@TableId(type = IdType.AUTO)
|
||||
private Long id;
|
||||
|
||||
private String modelCode;
|
||||
|
||||
private String displayName;
|
||||
|
||||
private String description;
|
||||
|
||||
private Boolean enabled;
|
||||
|
||||
private Boolean allowImage;
|
||||
|
||||
private Boolean allowVoice;
|
||||
|
||||
private Boolean isDefault;
|
||||
|
||||
private Integer sort;
|
||||
|
||||
@TableField(fill = FieldFill.INSERT)
|
||||
private LocalDateTime createdAt;
|
||||
|
||||
@TableField(fill = FieldFill.INSERT_UPDATE)
|
||||
private LocalDateTime updatedAt;
|
||||
|
||||
@TableLogic
|
||||
private Integer deleted;
|
||||
}
|
||||
@@ -14,6 +14,8 @@ public class AiChatSession {
|
||||
|
||||
private Long userId;
|
||||
|
||||
private Long modelId;
|
||||
|
||||
private String title;
|
||||
|
||||
private Integer roundCount;
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
package fun.nojava.module.blog.enums;
|
||||
|
||||
/**
|
||||
* AI 助手会话只读原因。
|
||||
*/
|
||||
public enum AiChatReadOnlyReason {
|
||||
ROUND_FULL,
|
||||
MODEL_DISABLED
|
||||
}
|
||||
+13
@@ -0,0 +1,13 @@
|
||||
package fun.nojava.module.blog.mapper;
|
||||
|
||||
import com.baomidou.mybatisplus.core.mapper.BaseMapper;
|
||||
import fun.nojava.module.blog.entity.AiAssistantModel;
|
||||
import org.apache.ibatis.annotations.Mapper;
|
||||
import org.apache.ibatis.annotations.Select;
|
||||
|
||||
@Mapper
|
||||
public interface AiAssistantModelMapper extends BaseMapper<AiAssistantModel> {
|
||||
|
||||
@Select("SELECT * FROM ai_assistant_model WHERE id = #{id}")
|
||||
AiAssistantModel selectByIdRaw(Long id);
|
||||
}
|
||||
+29
@@ -0,0 +1,29 @@
|
||||
package fun.nojava.module.blog.service;
|
||||
|
||||
import fun.nojava.module.blog.dto.AiAssistantModelAppVO;
|
||||
import fun.nojava.module.blog.dto.AiAssistantModelSaveDTO;
|
||||
import fun.nojava.module.blog.dto.AiAssistantModelVO;
|
||||
import fun.nojava.module.blog.entity.AiAssistantModel;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
public interface AiAssistantModelService {
|
||||
|
||||
List<AiAssistantModelVO> listAllForAdmin();
|
||||
|
||||
AiAssistantModelVO getByIdForAdmin(Long id);
|
||||
|
||||
AiAssistantModelVO create(AiAssistantModelSaveDTO dto);
|
||||
|
||||
void update(Long id, AiAssistantModelSaveDTO dto);
|
||||
|
||||
void deleteLogical(Long id);
|
||||
|
||||
List<AiAssistantModelAppVO> listEnabledForApp();
|
||||
|
||||
AiAssistantModel getByIdRaw(Long id);
|
||||
|
||||
AiAssistantModel requireEnabledForNewSession(Long modelId);
|
||||
|
||||
AiAssistantModel getDefaultModel();
|
||||
}
|
||||
@@ -7,11 +7,13 @@ import java.util.List;
|
||||
|
||||
public interface AiChatService {
|
||||
|
||||
CreateAiChatSessionVO createSession(Long userId);
|
||||
CreateAiChatSessionVO createSession(Long userId, Long modelId);
|
||||
|
||||
PageResult<AiChatSessionVO> listSessions(Long userId, int page, int size);
|
||||
|
||||
List<AiChatMessageVO> listMessages(Long userId, Long sessionId);
|
||||
AiChatSessionMessagesVO listMessages(Long userId, Long sessionId);
|
||||
|
||||
List<AiAssistantModelAppVO> listModelsForApp();
|
||||
|
||||
AiChatResultVO chat(Long userId, AiChatSendDTO dto);
|
||||
}
|
||||
|
||||
+239
@@ -0,0 +1,239 @@
|
||||
package fun.nojava.module.blog.service.impl;
|
||||
|
||||
import com.baomidou.mybatisplus.core.toolkit.Wrappers;
|
||||
import fun.nojava.common.exception.BusinessException;
|
||||
import fun.nojava.common.exception.ErrorCode;
|
||||
import fun.nojava.module.blog.config.DashScopeProperties;
|
||||
import fun.nojava.module.blog.dto.AiAssistantModelAppVO;
|
||||
import fun.nojava.module.blog.dto.AiAssistantModelSaveDTO;
|
||||
import fun.nojava.module.blog.dto.AiAssistantModelVO;
|
||||
import fun.nojava.module.blog.entity.AiAssistantModel;
|
||||
import fun.nojava.module.blog.mapper.AiAssistantModelMapper;
|
||||
import fun.nojava.module.blog.service.AiAssistantModelService;
|
||||
import lombok.RequiredArgsConstructor;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import org.springframework.stereotype.Service;
|
||||
import org.springframework.transaction.annotation.Transactional;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
@Slf4j
|
||||
@Service
|
||||
@RequiredArgsConstructor
|
||||
public class AiAssistantModelServiceImpl implements AiAssistantModelService {
|
||||
|
||||
private static final Pattern MODEL_CODE_PATTERN = Pattern.compile("^[a-zA-Z0-9._-]{1,64}$");
|
||||
|
||||
private final AiAssistantModelMapper modelMapper;
|
||||
private final DashScopeProperties dashScopeProperties;
|
||||
|
||||
@Override
|
||||
public List<AiAssistantModelVO> listAllForAdmin() {
|
||||
return modelMapper.selectList(Wrappers.<AiAssistantModel>lambdaQuery()
|
||||
.orderByAsc(AiAssistantModel::getSort)
|
||||
.orderByAsc(AiAssistantModel::getId))
|
||||
.stream()
|
||||
.map(this::toAdminVO)
|
||||
.toList();
|
||||
}
|
||||
|
||||
@Override
|
||||
public AiAssistantModelVO getByIdForAdmin(Long id) {
|
||||
AiAssistantModel model = modelMapper.selectById(id);
|
||||
if (model == null) {
|
||||
throw new BusinessException(ErrorCode.NOT_FOUND, "模型不存在");
|
||||
}
|
||||
return toAdminVO(model);
|
||||
}
|
||||
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public AiAssistantModelVO create(AiAssistantModelSaveDTO dto) {
|
||||
validateSaveDto(dto, null);
|
||||
AiAssistantModel model = fromSaveDto(dto);
|
||||
if (model.getEnabled() == null) {
|
||||
model.setEnabled(true);
|
||||
}
|
||||
if (model.getSort() == null) {
|
||||
model.setSort(0);
|
||||
}
|
||||
modelMapper.insert(model);
|
||||
if (Boolean.TRUE.equals(dto.getIsDefault())) {
|
||||
applyDefaultFlag(model.getId(), true);
|
||||
}
|
||||
return toAdminVO(modelMapper.selectById(model.getId()));
|
||||
}
|
||||
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void update(Long id, AiAssistantModelSaveDTO dto) {
|
||||
AiAssistantModel existing = modelMapper.selectById(id);
|
||||
if (existing == null) {
|
||||
throw new BusinessException(ErrorCode.NOT_FOUND, "模型不存在");
|
||||
}
|
||||
validateSaveDto(dto, id);
|
||||
|
||||
boolean wasEnabled = Boolean.TRUE.equals(existing.getEnabled());
|
||||
boolean willEnable = dto.getEnabled() == null ? wasEnabled : dto.getEnabled();
|
||||
if (wasEnabled && !willEnable) {
|
||||
assertAtLeastOneEnabledAfterChange(id);
|
||||
}
|
||||
|
||||
AiAssistantModel patch = fromSaveDto(dto);
|
||||
patch.setId(id);
|
||||
modelMapper.updateById(patch);
|
||||
|
||||
if (Boolean.TRUE.equals(dto.getIsDefault())) {
|
||||
applyDefaultFlag(id, true);
|
||||
} else if (Boolean.FALSE.equals(dto.getIsDefault()) && Boolean.TRUE.equals(existing.getIsDefault())) {
|
||||
throw new BusinessException(ErrorCode.AI_MODEL_DEFAULT_REQUIRED);
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
@Transactional(rollbackFor = Exception.class)
|
||||
public void deleteLogical(Long id) {
|
||||
AiAssistantModel model = modelMapper.selectById(id);
|
||||
if (model == null) {
|
||||
throw new BusinessException(ErrorCode.NOT_FOUND, "模型不存在");
|
||||
}
|
||||
if (Boolean.TRUE.equals(model.getIsDefault())) {
|
||||
throw new BusinessException(ErrorCode.AI_MODEL_DEFAULT_REQUIRED);
|
||||
}
|
||||
if (Boolean.TRUE.equals(model.getEnabled())) {
|
||||
assertAtLeastOneEnabledAfterChange(id);
|
||||
}
|
||||
modelMapper.deleteById(id);
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<AiAssistantModelAppVO> listEnabledForApp() {
|
||||
return modelMapper.selectList(Wrappers.<AiAssistantModel>lambdaQuery()
|
||||
.eq(AiAssistantModel::getEnabled, true)
|
||||
.orderByAsc(AiAssistantModel::getSort)
|
||||
.orderByAsc(AiAssistantModel::getId))
|
||||
.stream()
|
||||
.map(this::toAppVO)
|
||||
.toList();
|
||||
}
|
||||
|
||||
@Override
|
||||
public AiAssistantModel getByIdRaw(Long id) {
|
||||
if (id == null) {
|
||||
return null;
|
||||
}
|
||||
return modelMapper.selectByIdRaw(id);
|
||||
}
|
||||
|
||||
@Override
|
||||
public AiAssistantModel requireEnabledForNewSession(Long modelId) {
|
||||
if (modelId == null) {
|
||||
return getDefaultModel();
|
||||
}
|
||||
AiAssistantModel model = modelMapper.selectById(modelId);
|
||||
if (model == null || !Boolean.TRUE.equals(model.getEnabled())) {
|
||||
throw new BusinessException(ErrorCode.AI_MODEL_NOT_AVAILABLE);
|
||||
}
|
||||
return model;
|
||||
}
|
||||
|
||||
@Override
|
||||
public AiAssistantModel getDefaultModel() {
|
||||
AiAssistantModel model = modelMapper.selectOne(Wrappers.<AiAssistantModel>lambdaQuery()
|
||||
.eq(AiAssistantModel::getIsDefault, true)
|
||||
.last("LIMIT 1"));
|
||||
if (model != null) {
|
||||
return model;
|
||||
}
|
||||
log.warn("未配置默认 AI 助手模型,回退 dashscope.model={}", dashScopeProperties.getModel());
|
||||
AiAssistantModel fallback = new AiAssistantModel();
|
||||
fallback.setModelCode(dashScopeProperties.getModel());
|
||||
fallback.setDisplayName("默认模型");
|
||||
fallback.setEnabled(true);
|
||||
return fallback;
|
||||
}
|
||||
|
||||
private void validateSaveDto(AiAssistantModelSaveDTO dto, Long excludeId) {
|
||||
String code = dto.getModelCode() == null ? "" : dto.getModelCode().trim();
|
||||
if (!MODEL_CODE_PATTERN.matcher(code).matches()) {
|
||||
throw new BusinessException(ErrorCode.BAD_REQUEST, "模型编码格式不合法");
|
||||
}
|
||||
dto.setModelCode(code);
|
||||
dto.setDisplayName(dto.getDisplayName().trim());
|
||||
|
||||
long dup = modelMapper.selectCount(Wrappers.<AiAssistantModel>lambdaQuery()
|
||||
.eq(AiAssistantModel::getModelCode, code)
|
||||
.ne(excludeId != null, AiAssistantModel::getId, excludeId));
|
||||
if (dup > 0) {
|
||||
throw new BusinessException(ErrorCode.AI_MODEL_CODE_DUPLICATE);
|
||||
}
|
||||
}
|
||||
|
||||
private void assertAtLeastOneEnabledAfterChange(Long excludeId) {
|
||||
long enabledCount = modelMapper.selectCount(Wrappers.<AiAssistantModel>lambdaQuery()
|
||||
.eq(AiAssistantModel::getEnabled, true)
|
||||
.ne(AiAssistantModel::getId, excludeId));
|
||||
if (enabledCount == 0) {
|
||||
throw new BusinessException(ErrorCode.AI_MODEL_MUST_KEEP_ENABLED);
|
||||
}
|
||||
}
|
||||
|
||||
private void applyDefaultFlag(Long modelId, boolean isDefault) {
|
||||
if (!isDefault) {
|
||||
return;
|
||||
}
|
||||
modelMapper.update(null, Wrappers.<AiAssistantModel>lambdaUpdate()
|
||||
.set(AiAssistantModel::getIsDefault, false));
|
||||
AiAssistantModel patch = new AiAssistantModel();
|
||||
patch.setId(modelId);
|
||||
patch.setIsDefault(true);
|
||||
modelMapper.updateById(patch);
|
||||
}
|
||||
|
||||
private AiAssistantModel fromSaveDto(AiAssistantModelSaveDTO dto) {
|
||||
AiAssistantModel model = new AiAssistantModel();
|
||||
model.setModelCode(dto.getModelCode());
|
||||
model.setDisplayName(dto.getDisplayName());
|
||||
model.setDescription(dto.getDescription());
|
||||
if (dto.getEnabled() != null) {
|
||||
model.setEnabled(dto.getEnabled());
|
||||
}
|
||||
model.setAllowImage(Boolean.TRUE.equals(dto.getAllowImageUpload()));
|
||||
model.setAllowVoice(Boolean.TRUE.equals(dto.getAllowVoiceUpload()));
|
||||
if (dto.getIsDefault() != null) {
|
||||
model.setIsDefault(dto.getIsDefault());
|
||||
}
|
||||
if (dto.getSort() != null) {
|
||||
model.setSort(dto.getSort());
|
||||
}
|
||||
return model;
|
||||
}
|
||||
|
||||
private AiAssistantModelVO toAdminVO(AiAssistantModel model) {
|
||||
return AiAssistantModelVO.builder()
|
||||
.id(model.getId())
|
||||
.modelCode(model.getModelCode())
|
||||
.displayName(model.getDisplayName())
|
||||
.description(model.getDescription())
|
||||
.enabled(model.getEnabled())
|
||||
.allowImageUpload(Boolean.TRUE.equals(model.getAllowImage()))
|
||||
.allowVoiceUpload(Boolean.TRUE.equals(model.getAllowVoice()))
|
||||
.isDefault(model.getIsDefault())
|
||||
.sort(model.getSort())
|
||||
.createdAt(model.getCreatedAt())
|
||||
.updatedAt(model.getUpdatedAt())
|
||||
.build();
|
||||
}
|
||||
|
||||
private AiAssistantModelAppVO toAppVO(AiAssistantModel model) {
|
||||
return AiAssistantModelAppVO.builder()
|
||||
.id(model.getId())
|
||||
.displayName(model.getDisplayName())
|
||||
.description(model.getDescription())
|
||||
.isDefault(Boolean.TRUE.equals(model.getIsDefault()))
|
||||
.allowImageUpload(Boolean.TRUE.equals(model.getAllowImage()))
|
||||
.allowVoiceUpload(Boolean.TRUE.equals(model.getAllowVoice()))
|
||||
.build();
|
||||
}
|
||||
}
|
||||
+75
-12
@@ -8,10 +8,13 @@ import fun.nojava.common.exception.ErrorCode;
|
||||
import fun.nojava.common.model.PageResult;
|
||||
import fun.nojava.module.blog.config.DashScopeProperties;
|
||||
import fun.nojava.module.blog.dto.*;
|
||||
import fun.nojava.module.blog.entity.AiAssistantModel;
|
||||
import fun.nojava.module.blog.entity.AiChatMessage;
|
||||
import fun.nojava.module.blog.entity.AiChatSession;
|
||||
import fun.nojava.module.blog.enums.AiChatReadOnlyReason;
|
||||
import fun.nojava.module.blog.mapper.AiChatMessageMapper;
|
||||
import fun.nojava.module.blog.mapper.AiChatSessionMapper;
|
||||
import fun.nojava.module.blog.service.AiAssistantModelService;
|
||||
import fun.nojava.module.blog.service.AiChatService;
|
||||
import lombok.RequiredArgsConstructor;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
@@ -42,6 +45,7 @@ public class AiChatServiceImpl implements AiChatService {
|
||||
private static final String ROLE_USER = "user";
|
||||
private static final String ROLE_ASSISTANT = "assistant";
|
||||
private static final String REPLY_TRUNCATE_SUFFIX = "…";
|
||||
private static final String REMOVED_MODEL_LABEL = "已下架模型";
|
||||
|
||||
private static final String SYSTEM_PROMPT =
|
||||
"你是 Jog 博客系统的 AI 助手,友好、简洁地用中文回答用户问题。"
|
||||
@@ -54,11 +58,14 @@ public class AiChatServiceImpl implements AiChatService {
|
||||
private final AiSummaryServiceImpl aiSummaryService;
|
||||
private final DashScopeProperties properties;
|
||||
private final StringRedisTemplate redisTemplate;
|
||||
private final AiAssistantModelService aiAssistantModelService;
|
||||
|
||||
@Override
|
||||
public CreateAiChatSessionVO createSession(Long userId) {
|
||||
public CreateAiChatSessionVO createSession(Long userId, Long modelId) {
|
||||
AiAssistantModel model = aiAssistantModelService.requireEnabledForNewSession(modelId);
|
||||
AiChatSession session = new AiChatSession();
|
||||
session.setUserId(userId);
|
||||
session.setModelId(model.getId());
|
||||
session.setTitle("");
|
||||
session.setRoundCount(0);
|
||||
sessionMapper.insert(session);
|
||||
@@ -82,10 +89,24 @@ public class AiChatServiceImpl implements AiChatService {
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<AiChatMessageVO> listMessages(Long userId, Long sessionId) {
|
||||
getSessionForUser(sessionId, userId);
|
||||
public AiChatSessionMessagesVO listMessages(Long userId, Long sessionId) {
|
||||
AiChatSession session = getSessionForUser(sessionId, userId);
|
||||
BoundModelView bound = resolveBoundModel(session);
|
||||
List<AiChatMessage> messages = listMessagesBySession(sessionId);
|
||||
return messages.stream().map(this::toMessageVO).toList();
|
||||
AiChatReadOnlyReason reason = resolveReadOnlyReason(session, bound);
|
||||
|
||||
return AiChatSessionMessagesVO.builder()
|
||||
.messages(messages.stream().map(this::toMessageVO).toList())
|
||||
.modelEnabled(bound.enabled())
|
||||
.readOnlyReason(reason)
|
||||
.roundCount(session.getRoundCount())
|
||||
.maxRounds(MAX_ROUNDS)
|
||||
.build();
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<AiAssistantModelAppVO> listModelsForApp() {
|
||||
return aiAssistantModelService.listEnabledForApp();
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -101,6 +122,11 @@ public class AiChatServiceImpl implements AiChatService {
|
||||
throw new BusinessException(ErrorCode.AI_CHAT_SESSION_FULL);
|
||||
}
|
||||
|
||||
BoundModelView bound = resolveBoundModel(session);
|
||||
if (!bound.enabled()) {
|
||||
throw new BusinessException(ErrorCode.AI_MODEL_DISABLED);
|
||||
}
|
||||
|
||||
if (!hasApiKey()) {
|
||||
throw new BusinessException(ErrorCode.AI_API_KEY_MISSING);
|
||||
}
|
||||
@@ -128,7 +154,7 @@ public class AiChatServiceImpl implements AiChatService {
|
||||
}
|
||||
|
||||
List<AiChatMessage> history = listMessagesBySession(session.getId());
|
||||
Map<String, Object> requestBody = buildRequestBody(history);
|
||||
Map<String, Object> requestBody = buildRequestBody(history, bound.modelCode());
|
||||
long startMillis = System.currentTimeMillis();
|
||||
|
||||
Map<String, Object> response;
|
||||
@@ -136,7 +162,8 @@ public class AiChatServiceImpl implements AiChatService {
|
||||
response = aiSummaryService.callModel(requestBody);
|
||||
} catch (ResourceAccessException ex) {
|
||||
if (ex.getCause() instanceof SocketTimeoutException) {
|
||||
log.error("AI 助手调用超时, userId={}, sessionId={}", userId, session.getId());
|
||||
log.error("AI 助手调用超时, userId={}, sessionId={}, modelId={}",
|
||||
userId, session.getId(), session.getModelId());
|
||||
throw new BusinessException(ErrorCode.AI_SERVICE_TIMEOUT);
|
||||
}
|
||||
log.error("AI 助手网络异常, userId={}, sessionId={}, error={}", userId, session.getId(), ex.getMessage());
|
||||
@@ -174,23 +201,55 @@ public class AiChatServiceImpl implements AiChatService {
|
||||
promptTokens = toInt(usage.get("prompt_tokens"));
|
||||
completionTokens = toInt(usage.get("completion_tokens"));
|
||||
}
|
||||
String model = (String) response.getOrDefault("model", properties.getModel());
|
||||
String modelCode = (String) response.getOrDefault("model", bound.modelCode());
|
||||
|
||||
log.info("AI 助手回复成功 userId={}, sessionId={}, roundCount={}, replyLen={}, cost={}ms",
|
||||
userId, session.getId(), newRoundCount, reply.length(), costMillis);
|
||||
log.info("AI 助手回复成功 userId={}, sessionId={}, modelId={}, modelCode={}, roundCount={}, replyLen={}, cost={}ms",
|
||||
userId, session.getId(), session.getModelId(), modelCode, newRoundCount, reply.length(), costMillis);
|
||||
|
||||
return AiChatResultVO.builder()
|
||||
.sessionId(session.getId())
|
||||
.reply(reply)
|
||||
.roundCount(newRoundCount)
|
||||
.maxRounds(MAX_ROUNDS)
|
||||
.model(model)
|
||||
.model(modelCode)
|
||||
.modelDisplayName(bound.displayName())
|
||||
.promptTokens(promptTokens)
|
||||
.completionTokens(completionTokens)
|
||||
.costMillis(costMillis)
|
||||
.build();
|
||||
}
|
||||
|
||||
private record BoundModelView(String displayName, boolean enabled, String modelCode) {}
|
||||
|
||||
private BoundModelView resolveBoundModel(AiChatSession session) {
|
||||
Long modelId = session.getModelId();
|
||||
if (modelId == null) {
|
||||
AiAssistantModel fallback = aiAssistantModelService.getDefaultModel();
|
||||
return new BoundModelView(
|
||||
fallback.getDisplayName(),
|
||||
Boolean.TRUE.equals(fallback.getEnabled()),
|
||||
fallback.getModelCode());
|
||||
}
|
||||
AiAssistantModel model = aiAssistantModelService.getByIdRaw(modelId);
|
||||
if (model == null) {
|
||||
return new BoundModelView(REMOVED_MODEL_LABEL, false, properties.getModel());
|
||||
}
|
||||
return new BoundModelView(
|
||||
model.getDisplayName(),
|
||||
Boolean.TRUE.equals(model.getEnabled()),
|
||||
model.getModelCode());
|
||||
}
|
||||
|
||||
private AiChatReadOnlyReason resolveReadOnlyReason(AiChatSession session, BoundModelView bound) {
|
||||
if (!bound.enabled()) {
|
||||
return AiChatReadOnlyReason.MODEL_DISABLED;
|
||||
}
|
||||
if (session.getRoundCount() != null && session.getRoundCount() >= MAX_ROUNDS) {
|
||||
return AiChatReadOnlyReason.ROUND_FULL;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
private AiChatSession getSessionForUser(Long sessionId, Long userId) {
|
||||
AiChatSession session = sessionMapper.selectById(sessionId);
|
||||
if (session == null || !userId.equals(session.getUserId())) {
|
||||
@@ -217,7 +276,7 @@ public class AiChatServiceImpl implements AiChatService {
|
||||
return last == null ? 1 : last.getSortOrder() + 1;
|
||||
}
|
||||
|
||||
private Map<String, Object> buildRequestBody(List<AiChatMessage> history) {
|
||||
private Map<String, Object> buildRequestBody(List<AiChatMessage> history, String modelCode) {
|
||||
List<Map<String, String>> messages = new ArrayList<>();
|
||||
messages.add(Map.of("role", "system", "content", SYSTEM_PROMPT));
|
||||
|
||||
@@ -227,7 +286,7 @@ public class AiChatServiceImpl implements AiChatService {
|
||||
}
|
||||
|
||||
return Map.of(
|
||||
"model", properties.getModel(),
|
||||
"model", modelCode,
|
||||
"messages", messages,
|
||||
"temperature", 0.7,
|
||||
"top_p", 0.9
|
||||
@@ -297,11 +356,15 @@ public class AiChatServiceImpl implements AiChatService {
|
||||
}
|
||||
|
||||
private AiChatSessionVO toSessionVO(AiChatSession session) {
|
||||
BoundModelView bound = resolveBoundModel(session);
|
||||
return AiChatSessionVO.builder()
|
||||
.id(session.getId())
|
||||
.title(session.getTitle() == null || session.getTitle().isBlank() ? "新会话" : session.getTitle())
|
||||
.roundCount(session.getRoundCount())
|
||||
.lastMessageAt(session.getLastMessageAt())
|
||||
.modelId(session.getModelId())
|
||||
.modelDisplayName(bound.displayName())
|
||||
.modelEnabled(bound.enabled())
|
||||
.build();
|
||||
}
|
||||
|
||||
|
||||
+159
@@ -0,0 +1,159 @@
|
||||
package fun.nojava.module.blog.service.impl;
|
||||
|
||||
import com.baomidou.mybatisplus.core.MybatisConfiguration;
|
||||
import com.baomidou.mybatisplus.core.conditions.Wrapper;
|
||||
import com.baomidou.mybatisplus.core.metadata.TableInfoHelper;
|
||||
import fun.nojava.common.exception.BusinessException;
|
||||
import fun.nojava.common.exception.ErrorCode;
|
||||
import fun.nojava.module.blog.config.DashScopeProperties;
|
||||
import fun.nojava.module.blog.dto.AiAssistantModelSaveDTO;
|
||||
import fun.nojava.module.blog.entity.AiAssistantModel;
|
||||
import fun.nojava.module.blog.mapper.AiAssistantModelMapper;
|
||||
import org.apache.ibatis.builder.MapperBuilderAssistant;
|
||||
import org.junit.jupiter.api.BeforeEach;
|
||||
import org.junit.jupiter.api.DisplayName;
|
||||
import org.junit.jupiter.api.Tag;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.junit.jupiter.api.extension.ExtendWith;
|
||||
import org.mockito.ArgumentCaptor;
|
||||
import org.mockito.Mock;
|
||||
import org.mockito.junit.jupiter.MockitoExtension;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
import static org.mockito.ArgumentMatchers.*;
|
||||
import static org.mockito.Mockito.*;
|
||||
|
||||
@ExtendWith(MockitoExtension.class)
|
||||
class AiAssistantModelServiceImplTest {
|
||||
|
||||
@Mock
|
||||
private AiAssistantModelMapper modelMapper;
|
||||
|
||||
private AiAssistantModelServiceImpl service;
|
||||
|
||||
@BeforeEach
|
||||
void setUp() {
|
||||
TableInfoHelper.initTableInfo(
|
||||
new MapperBuilderAssistant(new MybatisConfiguration(), ""),
|
||||
AiAssistantModel.class);
|
||||
DashScopeProperties properties = new DashScopeProperties();
|
||||
properties.setModel("qwen-turbo");
|
||||
service = new AiAssistantModelServiceImpl(modelMapper, properties);
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p0")
|
||||
@DisplayName("TC-AI-MODEL-01: 重复 modelCode 应抛出 AI_MODEL_CODE_DUPLICATE")
|
||||
void create_shouldRejectDuplicateModelCode() {
|
||||
when(modelMapper.selectCount(any())).thenReturn(1L);
|
||||
|
||||
AiAssistantModelSaveDTO dto = new AiAssistantModelSaveDTO();
|
||||
dto.setModelCode("qwen-turbo");
|
||||
dto.setDisplayName("通义 Turbo");
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class, () -> service.create(dto));
|
||||
assertEquals(ErrorCode.AI_MODEL_CODE_DUPLICATE.getCode(), ex.getCode());
|
||||
verify(modelMapper, never()).insert(any(AiAssistantModel.class));
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p0")
|
||||
@DisplayName("TC-AI-MODEL-02: 设为默认模型应清除其它默认标记")
|
||||
void create_shouldApplyDefaultFlagExclusively() {
|
||||
when(modelMapper.selectCount(any())).thenReturn(0L);
|
||||
doAnswer(inv -> {
|
||||
AiAssistantModel m = inv.getArgument(0);
|
||||
m.setId(2L);
|
||||
return 1;
|
||||
}).when(modelMapper).insert(any(AiAssistantModel.class));
|
||||
|
||||
AiAssistantModel saved = new AiAssistantModel();
|
||||
saved.setId(2L);
|
||||
saved.setModelCode("qwen-plus");
|
||||
saved.setDisplayName("通义 Plus");
|
||||
saved.setEnabled(true);
|
||||
saved.setIsDefault(true);
|
||||
when(modelMapper.selectById(2L)).thenReturn(saved);
|
||||
when(modelMapper.update(isNull(), any())).thenReturn(1);
|
||||
when(modelMapper.updateById(any(AiAssistantModel.class))).thenReturn(1);
|
||||
|
||||
AiAssistantModelSaveDTO dto = new AiAssistantModelSaveDTO();
|
||||
dto.setModelCode("qwen-plus");
|
||||
dto.setDisplayName("通义 Plus");
|
||||
dto.setIsDefault(true);
|
||||
|
||||
service.create(dto);
|
||||
|
||||
verify(modelMapper).update(isNull(), any(Wrapper.class));
|
||||
ArgumentCaptor<AiAssistantModel> patchCaptor = ArgumentCaptor.forClass(AiAssistantModel.class);
|
||||
verify(modelMapper).updateById(patchCaptor.capture());
|
||||
assertTrue(Boolean.TRUE.equals(patchCaptor.getValue().getIsDefault()));
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p0")
|
||||
@DisplayName("TC-AI-MODEL-03: 禁用最后一个启用模型应抛出 AI_MODEL_MUST_KEEP_ENABLED")
|
||||
void update_shouldRejectDisablingLastEnabledModel() {
|
||||
AiAssistantModel existing = new AiAssistantModel();
|
||||
existing.setId(1L);
|
||||
existing.setModelCode("qwen-turbo");
|
||||
existing.setDisplayName("默认模型");
|
||||
existing.setEnabled(true);
|
||||
when(modelMapper.selectById(1L)).thenReturn(existing);
|
||||
when(modelMapper.selectCount(any())).thenReturn(0L, 0L);
|
||||
|
||||
AiAssistantModelSaveDTO dto = new AiAssistantModelSaveDTO();
|
||||
dto.setModelCode("qwen-turbo");
|
||||
dto.setDisplayName("默认模型");
|
||||
dto.setEnabled(false);
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class, () -> service.update(1L, dto));
|
||||
assertEquals(ErrorCode.AI_MODEL_MUST_KEEP_ENABLED.getCode(), ex.getCode());
|
||||
verify(modelMapper, never()).updateById(any(AiAssistantModel.class));
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p0")
|
||||
@DisplayName("TC-AI-MODEL-04: 删除默认模型应抛出 AI_MODEL_DEFAULT_REQUIRED")
|
||||
void delete_shouldRejectDefaultModel() {
|
||||
AiAssistantModel model = new AiAssistantModel();
|
||||
model.setId(1L);
|
||||
model.setModelCode("qwen-turbo");
|
||||
model.setDisplayName("默认模型");
|
||||
model.setEnabled(true);
|
||||
model.setIsDefault(true);
|
||||
when(modelMapper.selectById(1L)).thenReturn(model);
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class, () -> service.deleteLogical(1L));
|
||||
assertEquals(ErrorCode.AI_MODEL_DEFAULT_REQUIRED.getCode(), ex.getCode());
|
||||
verify(modelMapper, never()).deleteById(anyLong());
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p1")
|
||||
@DisplayName("TC-AI-MODEL-05: 创建会话指定无效 modelId 应抛出 AI_MODEL_NOT_AVAILABLE")
|
||||
void requireEnabledForNewSession_shouldRejectInvalidModelId() {
|
||||
when(modelMapper.selectById(99L)).thenReturn(null);
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class,
|
||||
() -> service.requireEnabledForNewSession(99L));
|
||||
assertEquals(ErrorCode.AI_MODEL_NOT_AVAILABLE.getCode(), ex.getCode());
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p1")
|
||||
@DisplayName("TC-AI-MODEL-06: 未传 modelId 创建会话应使用默认模型")
|
||||
void requireEnabledForNewSession_shouldUseDefaultWhenNull() {
|
||||
AiAssistantModel defaultModel = new AiAssistantModel();
|
||||
defaultModel.setId(1L);
|
||||
defaultModel.setModelCode("qwen-turbo");
|
||||
defaultModel.setDisplayName("默认模型");
|
||||
defaultModel.setEnabled(true);
|
||||
when(modelMapper.selectOne(any())).thenReturn(defaultModel);
|
||||
|
||||
AiAssistantModel result = service.requireEnabledForNewSession(null);
|
||||
|
||||
assertEquals(1L, result.getId());
|
||||
assertEquals("qwen-turbo", result.getModelCode());
|
||||
}
|
||||
}
|
||||
+70
@@ -0,0 +1,70 @@
|
||||
package fun.nojava.module.blog.service.impl;
|
||||
|
||||
import fun.nojava.module.blog.entity.AiAssistantModel;
|
||||
import fun.nojava.module.blog.service.AiAssistantModelService;
|
||||
|
||||
import static org.mockito.ArgumentMatchers.any;
|
||||
import static org.mockito.ArgumentMatchers.anyLong;
|
||||
import static org.mockito.ArgumentMatchers.isNull;
|
||||
import static org.mockito.Mockito.lenient;
|
||||
import static org.mockito.Mockito.when;
|
||||
|
||||
/**
|
||||
* 为 AI Chat 测试提供一致的默认模型 Stub。
|
||||
*/
|
||||
final class AiAssistantModelTestSupport {
|
||||
|
||||
static final long DEFAULT_MODEL_ID = 1L;
|
||||
static final String DEFAULT_MODEL_CODE = "qwen-turbo";
|
||||
static final String DEFAULT_DISPLAY_NAME = "默认模型";
|
||||
|
||||
private AiAssistantModelTestSupport() {
|
||||
}
|
||||
|
||||
static AiAssistantModel defaultModel() {
|
||||
AiAssistantModel model = new AiAssistantModel();
|
||||
model.setId(DEFAULT_MODEL_ID);
|
||||
model.setModelCode(DEFAULT_MODEL_CODE);
|
||||
model.setDisplayName(DEFAULT_DISPLAY_NAME);
|
||||
model.setEnabled(true);
|
||||
model.setIsDefault(true);
|
||||
model.setDeleted(0);
|
||||
return model;
|
||||
}
|
||||
|
||||
static AiAssistantModel model(long id, String code, String displayName, boolean enabled) {
|
||||
AiAssistantModel model = new AiAssistantModel();
|
||||
model.setId(id);
|
||||
model.setModelCode(code);
|
||||
model.setDisplayName(displayName);
|
||||
model.setEnabled(enabled);
|
||||
model.setDeleted(0);
|
||||
return model;
|
||||
}
|
||||
|
||||
static void wireDefaultModel(AiAssistantModelService modelService) {
|
||||
AiAssistantModel model = defaultModel();
|
||||
lenient().when(modelService.getDefaultModel()).thenReturn(model);
|
||||
lenient().when(modelService.requireEnabledForNewSession(isNull())).thenReturn(model);
|
||||
lenient().when(modelService.requireEnabledForNewSession(anyLong())).thenAnswer(inv -> {
|
||||
Long modelId = inv.getArgument(0);
|
||||
if (modelId == null || modelId.equals(DEFAULT_MODEL_ID)) {
|
||||
return model;
|
||||
}
|
||||
AiAssistantModel other = model(modelId, "qwen-plus", "通义 Plus", true);
|
||||
return other;
|
||||
});
|
||||
lenient().when(modelService.getByIdRaw(isNull())).thenReturn(null);
|
||||
lenient().when(modelService.getByIdRaw(anyLong())).thenAnswer(inv -> {
|
||||
Long id = inv.getArgument(0);
|
||||
if (id == null) {
|
||||
return null;
|
||||
}
|
||||
if (id.equals(DEFAULT_MODEL_ID)) {
|
||||
return model;
|
||||
}
|
||||
return model(id, "qwen-plus", "通义 Plus", true);
|
||||
});
|
||||
lenient().when(modelService.listEnabledForApp()).thenReturn(java.util.List.of());
|
||||
}
|
||||
}
|
||||
+112
-49
@@ -3,23 +3,25 @@ package fun.nojava.module.blog.service.impl;
|
||||
import com.baomidou.mybatisplus.core.conditions.query.LambdaQueryWrapper;
|
||||
import fun.nojava.common.exception.BusinessException;
|
||||
import fun.nojava.common.exception.ErrorCode;
|
||||
import fun.nojava.module.blog.config.DashScopeProperties;
|
||||
import fun.nojava.module.blog.dto.AiChatSendDTO;
|
||||
import fun.nojava.module.blog.dto.AiChatResultVO;
|
||||
import fun.nojava.module.blog.dto.AiChatSendDTO;
|
||||
import fun.nojava.module.blog.dto.AiChatSessionMessagesVO;
|
||||
import fun.nojava.module.blog.dto.CreateAiChatSessionVO;
|
||||
import fun.nojava.module.blog.entity.AiAssistantModel;
|
||||
import fun.nojava.module.blog.entity.AiChatMessage;
|
||||
import fun.nojava.module.blog.entity.AiChatSession;
|
||||
import fun.nojava.module.blog.enums.AiChatReadOnlyReason;
|
||||
import fun.nojava.module.blog.mapper.AiChatMessageMapper;
|
||||
import fun.nojava.module.blog.mapper.AiChatSessionMapper;
|
||||
import fun.nojava.module.blog.service.AiAssistantModelService;
|
||||
import fun.nojava.module.blog.testutil.BaseAiServiceTest;
|
||||
import fun.nojava.module.blog.testutil.TestDataFactory;
|
||||
import org.junit.jupiter.api.BeforeEach;
|
||||
import org.junit.jupiter.api.DisplayName;
|
||||
import org.junit.jupiter.api.Tag;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.junit.jupiter.api.extension.ExtendWith;
|
||||
import org.mockito.ArgumentCaptor;
|
||||
import org.mockito.Mock;
|
||||
import org.mockito.junit.jupiter.MockitoExtension;
|
||||
import org.springframework.data.redis.core.StringRedisTemplate;
|
||||
import org.springframework.data.redis.core.ValueOperations;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
@@ -29,8 +31,7 @@ import static org.junit.jupiter.api.Assertions.*;
|
||||
import static org.mockito.ArgumentMatchers.*;
|
||||
import static org.mockito.Mockito.*;
|
||||
|
||||
@ExtendWith(MockitoExtension.class)
|
||||
class AiChatServiceImplTest {
|
||||
class AiChatServiceImplTest extends BaseAiServiceTest {
|
||||
|
||||
@Mock
|
||||
private AiChatSessionMapper sessionMapper;
|
||||
@@ -39,41 +40,33 @@ class AiChatServiceImplTest {
|
||||
@Mock
|
||||
private AiSummaryServiceImpl aiSummaryService;
|
||||
@Mock
|
||||
private StringRedisTemplate redisTemplate;
|
||||
@Mock
|
||||
private ValueOperations<String, String> valueOps;
|
||||
private AiAssistantModelService aiAssistantModelService;
|
||||
|
||||
private AiChatServiceImpl service;
|
||||
|
||||
private DashScopeProperties properties;
|
||||
|
||||
@BeforeEach
|
||||
void setUp() {
|
||||
properties = new DashScopeProperties();
|
||||
properties.setApiKey("test-key");
|
||||
properties.setModel("qwen-turbo");
|
||||
properties.setMaxInput(8000);
|
||||
service = new AiChatServiceImpl(sessionMapper, messageMapper, aiSummaryService, properties, redisTemplate);
|
||||
}
|
||||
|
||||
@SuppressWarnings("unchecked")
|
||||
private void stubRateLimitOk() {
|
||||
when(redisTemplate.opsForValue()).thenReturn(valueOps);
|
||||
when(valueOps.increment(anyString())).thenReturn(1L);
|
||||
service = new AiChatServiceImpl(
|
||||
sessionMapper, messageMapper, aiSummaryService, dashProps, redisTemplate, aiAssistantModelService);
|
||||
AiAssistantModelTestSupport.wireDefaultModel(aiAssistantModelService);
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p0")
|
||||
@DisplayName("创建会话应返回 sessionId")
|
||||
void createSession_shouldReturnId() {
|
||||
@DisplayName("创建会话应返回 sessionId 并绑定 modelId")
|
||||
void createSession_shouldReturnIdAndBindModel() {
|
||||
doAnswer(inv -> {
|
||||
AiChatSession s = inv.getArgument(0);
|
||||
s.setId(100L);
|
||||
return 1;
|
||||
}).when(sessionMapper).insert(any(AiChatSession.class));
|
||||
|
||||
CreateAiChatSessionVO vo = service.createSession(1L);
|
||||
CreateAiChatSessionVO vo = service.createSession(1L, null);
|
||||
assertEquals(100L, vo.getSessionId());
|
||||
|
||||
ArgumentCaptor<AiChatSession> captor = ArgumentCaptor.forClass(AiChatSession.class);
|
||||
verify(sessionMapper).insert(captor.capture());
|
||||
assertEquals(AiAssistantModelTestSupport.DEFAULT_MODEL_ID, captor.getValue().getModelId());
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -83,12 +76,11 @@ class AiChatServiceImplTest {
|
||||
AiChatSession session = new AiChatSession();
|
||||
session.setId(1L);
|
||||
session.setUserId(1L);
|
||||
session.setModelId(AiAssistantModelTestSupport.DEFAULT_MODEL_ID);
|
||||
session.setRoundCount(5);
|
||||
when(sessionMapper.selectById(1L)).thenReturn(session);
|
||||
|
||||
AiChatSendDTO dto = new AiChatSendDTO();
|
||||
dto.setSessionId(1L);
|
||||
dto.setMessage("继续问一个问题看看会怎样");
|
||||
AiChatSendDTO dto = TestDataFactory.aChatSendDTO(1L, "继续问一个问题看看会怎样");
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class, () -> service.chat(1L, dto));
|
||||
assertEquals(ErrorCode.AI_CHAT_SESSION_FULL.getCode(), ex.getCode());
|
||||
@@ -102,26 +94,46 @@ class AiChatServiceImplTest {
|
||||
AiChatSession session = new AiChatSession();
|
||||
session.setId(1L);
|
||||
session.setUserId(2L);
|
||||
session.setModelId(AiAssistantModelTestSupport.DEFAULT_MODEL_ID);
|
||||
session.setRoundCount(0);
|
||||
when(sessionMapper.selectById(1L)).thenReturn(session);
|
||||
|
||||
AiChatSendDTO dto = new AiChatSendDTO();
|
||||
dto.setSessionId(1L);
|
||||
dto.setMessage("这是一条足够长的测试问题用于 AI 助手单元测试");
|
||||
AiChatSendDTO dto = TestDataFactory.aChatSendDTO(1L, "这是一条足够长的测试问题");
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class, () -> service.chat(1L, dto));
|
||||
assertEquals(ErrorCode.NOT_FOUND.getCode(), ex.getCode());
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p0")
|
||||
@DisplayName("绑定模型已禁用应抛出 AI_MODEL_DISABLED")
|
||||
void chat_shouldRejectWhenBoundModelDisabled() {
|
||||
AiAssistantModel disabled = AiAssistantModelTestSupport.defaultModel();
|
||||
disabled.setEnabled(false);
|
||||
when(aiAssistantModelService.getByIdRaw(AiAssistantModelTestSupport.DEFAULT_MODEL_ID)).thenReturn(disabled);
|
||||
|
||||
AiChatSession session = new AiChatSession();
|
||||
session.setId(1L);
|
||||
session.setUserId(1L);
|
||||
session.setModelId(AiAssistantModelTestSupport.DEFAULT_MODEL_ID);
|
||||
session.setRoundCount(0);
|
||||
when(sessionMapper.selectById(1L)).thenReturn(session);
|
||||
|
||||
AiChatSendDTO dto = TestDataFactory.aChatSendDTO(1L, "这是一条足够长的测试问题");
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class, () -> service.chat(1L, dto));
|
||||
assertEquals(ErrorCode.AI_MODEL_DISABLED.getCode(), ex.getCode());
|
||||
verify(aiSummaryService, never()).callModel(any());
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p0")
|
||||
@DisplayName("首条消息成功应落库 user/assistant 且 roundCount=1")
|
||||
void chat_shouldSucceedOnFirstMessage() throws Exception {
|
||||
stubRateLimitOk();
|
||||
|
||||
AiChatSession session = new AiChatSession();
|
||||
session.setId(10L);
|
||||
session.setUserId(1L);
|
||||
session.setModelId(AiAssistantModelTestSupport.DEFAULT_MODEL_ID);
|
||||
session.setRoundCount(0);
|
||||
session.setTitle("");
|
||||
when(sessionMapper.selectById(10L)).thenReturn(session);
|
||||
@@ -133,16 +145,10 @@ class AiChatServiceImplTest {
|
||||
when(messageMapper.selectOne(any())).thenReturn(null, userMsg);
|
||||
when(messageMapper.selectList(any(LambdaQueryWrapper.class))).thenReturn(List.of(userMsg));
|
||||
|
||||
Map<String, Object> modelResponse = Map.of(
|
||||
"choices", List.of(Map.of("message", Map.of("content", "REST 是一种架构风格……"))),
|
||||
"model", "qwen-turbo",
|
||||
"usage", Map.of("prompt_tokens", 10, "completion_tokens", 20)
|
||||
);
|
||||
when(aiSummaryService.callModel(any())).thenReturn(modelResponse);
|
||||
when(aiSummaryService.callModel(any()))
|
||||
.thenReturn(TestDataFactory.modelResponse("REST 是一种架构风格……"));
|
||||
|
||||
AiChatSendDTO dto = new AiChatSendDTO();
|
||||
dto.setSessionId(10L);
|
||||
dto.setMessage("什么是 REST API?请简要说明。");
|
||||
AiChatSendDTO dto = TestDataFactory.aChatSendDTO(10L, "什么是 REST API?请简要说明。");
|
||||
|
||||
AiChatResultVO result = service.chat(1L, dto);
|
||||
|
||||
@@ -152,22 +158,79 @@ class AiChatServiceImplTest {
|
||||
verify(sessionMapper, atLeastOnce()).updateById(any(AiChatSession.class));
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p0")
|
||||
@DisplayName("chat 应使用会话绑定模型的 modelCode 调用 DashScope")
|
||||
void chat_shouldUseBoundModelCode() throws Exception {
|
||||
long plusModelId = 2L;
|
||||
AiAssistantModel plus = AiAssistantModelTestSupport.model(plusModelId, "qwen-plus", "通义 Plus", true);
|
||||
when(aiAssistantModelService.getByIdRaw(plusModelId)).thenReturn(plus);
|
||||
|
||||
AiChatSession session = new AiChatSession();
|
||||
session.setId(10L);
|
||||
session.setUserId(1L);
|
||||
session.setModelId(plusModelId);
|
||||
session.setRoundCount(0);
|
||||
session.setTitle("");
|
||||
when(sessionMapper.selectById(10L)).thenReturn(session);
|
||||
|
||||
AiChatMessage userMsg = new AiChatMessage();
|
||||
userMsg.setRole("user");
|
||||
userMsg.setContent("什么是 REST API?请简要说明。");
|
||||
userMsg.setSortOrder(1);
|
||||
when(messageMapper.selectOne(any())).thenReturn(null, userMsg);
|
||||
when(messageMapper.selectList(any(LambdaQueryWrapper.class))).thenReturn(List.of(userMsg));
|
||||
|
||||
when(aiSummaryService.callModel(any())).thenAnswer(inv -> {
|
||||
@SuppressWarnings("unchecked")
|
||||
Map<String, Object> body = inv.getArgument(0);
|
||||
assertEquals("qwen-plus", body.get("model"));
|
||||
return TestDataFactory.modelResponse("Plus 模型回复", "qwen-plus");
|
||||
});
|
||||
|
||||
AiChatSendDTO dto = TestDataFactory.aChatSendDTO(10L, "什么是 REST API?请简要说明。");
|
||||
|
||||
AiChatResultVO result = service.chat(1L, dto);
|
||||
assertEquals("qwen-plus", result.getModel());
|
||||
assertEquals("通义 Plus", result.getModelDisplayName());
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p0")
|
||||
@DisplayName("listMessages 绑定模型禁用时 readOnlyReason 应为 MODEL_DISABLED")
|
||||
void listMessages_shouldReturnModelDisabledReason() {
|
||||
AiAssistantModel disabled = AiAssistantModelTestSupport.defaultModel();
|
||||
disabled.setEnabled(false);
|
||||
when(aiAssistantModelService.getByIdRaw(AiAssistantModelTestSupport.DEFAULT_MODEL_ID)).thenReturn(disabled);
|
||||
|
||||
AiChatSession session = new AiChatSession();
|
||||
session.setId(1L);
|
||||
session.setUserId(1L);
|
||||
session.setModelId(AiAssistantModelTestSupport.DEFAULT_MODEL_ID);
|
||||
session.setRoundCount(2);
|
||||
when(sessionMapper.selectById(1L)).thenReturn(session);
|
||||
when(messageMapper.selectList(any(LambdaQueryWrapper.class))).thenReturn(List.of());
|
||||
|
||||
AiChatSessionMessagesVO vo = service.listMessages(1L, 1L);
|
||||
|
||||
assertFalse(vo.getModelEnabled());
|
||||
assertEquals(AiChatReadOnlyReason.MODEL_DISABLED, vo.getReadOnlyReason());
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p1")
|
||||
@DisplayName("限流超过 10 次应抛出 RATE_LIMITED")
|
||||
void chat_shouldRateLimit() {
|
||||
stubRateLimitOk();
|
||||
when(valueOps.increment("ai_chat:rate:1")).thenReturn(11L);
|
||||
|
||||
AiChatSession session = new AiChatSession();
|
||||
session.setId(1L);
|
||||
session.setUserId(1L);
|
||||
session.setModelId(AiAssistantModelTestSupport.DEFAULT_MODEL_ID);
|
||||
session.setRoundCount(0);
|
||||
when(sessionMapper.selectById(1L)).thenReturn(session);
|
||||
|
||||
AiChatSendDTO dto = new AiChatSendDTO();
|
||||
dto.setSessionId(1L);
|
||||
dto.setMessage("这是一条足够长的测试问题用于 AI 助手单元测试");
|
||||
AiChatSendDTO dto = TestDataFactory.aChatSendDTO(1L, "这是一条足够长的测试问题");
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class, () -> service.chat(1L, dto));
|
||||
assertEquals(ErrorCode.RATE_LIMITED.getCode(), ex.getCode());
|
||||
@@ -178,7 +241,7 @@ class AiChatServiceImplTest {
|
||||
@Tag("p1")
|
||||
@DisplayName("trimHistoryByMaxInput 应从最早消息丢弃")
|
||||
void trimHistoryByMaxInput_shouldDropOldest() {
|
||||
properties.setMaxInput(50);
|
||||
dashProps.setMaxInput(50);
|
||||
List<AiChatMessage> history = new ArrayList<>();
|
||||
for (int i = 0; i < 4; i++) {
|
||||
AiChatMessage u = new AiChatMessage();
|
||||
|
||||
+19
-12
@@ -7,6 +7,7 @@ import fun.nojava.module.blog.config.DashScopeProperties;
|
||||
import fun.nojava.module.blog.dto.*;
|
||||
import fun.nojava.module.blog.mapper.AiChatMessageMapper;
|
||||
import fun.nojava.module.blog.mapper.AiChatSessionMapper;
|
||||
import fun.nojava.module.blog.service.AiAssistantModelService;
|
||||
import org.junit.jupiter.api.AfterEach;
|
||||
import org.junit.jupiter.api.BeforeEach;
|
||||
import org.junit.jupiter.api.DisplayName;
|
||||
@@ -45,6 +46,8 @@ class AiChatServiceIntegrationTest {
|
||||
private StringRedisTemplate redisTemplate;
|
||||
@Mock
|
||||
private ValueOperations<String, String> valueOps;
|
||||
@Mock
|
||||
private AiAssistantModelService aiAssistantModelService;
|
||||
|
||||
private final AiChatInMemoryTestSupport store = new AiChatInMemoryTestSupport();
|
||||
private AiChatServiceImpl service;
|
||||
@@ -59,7 +62,9 @@ class AiChatServiceIntegrationTest {
|
||||
properties.setModel("qwen-turbo");
|
||||
properties.setMaxInput(8000);
|
||||
|
||||
service = new AiChatServiceImpl(sessionMapper, messageMapper, aiSummaryService, properties, redisTemplate);
|
||||
service = new AiChatServiceImpl(
|
||||
sessionMapper, messageMapper, aiSummaryService, properties, redisTemplate, aiAssistantModelService);
|
||||
AiAssistantModelTestSupport.wireDefaultModel(aiAssistantModelService);
|
||||
lenient().when(redisTemplate.opsForValue()).thenReturn(valueOps);
|
||||
lenient().when(valueOps.increment(any())).thenReturn(1L);
|
||||
}
|
||||
@@ -75,7 +80,7 @@ class AiChatServiceIntegrationTest {
|
||||
void shouldCompleteThreeRoundConversation() throws Exception {
|
||||
stubModelReply("回复一", "回复二", "回复三");
|
||||
|
||||
CreateAiChatSessionVO created = service.createSession(USER_ID);
|
||||
CreateAiChatSessionVO created = service.createSession(USER_ID, null);
|
||||
Long sessionId = created.getSessionId();
|
||||
|
||||
chat(sessionId, "问题一:什么是 REST?");
|
||||
@@ -85,17 +90,20 @@ class AiChatServiceIntegrationTest {
|
||||
assertEquals(3, r3.getRoundCount());
|
||||
assertEquals(5, r3.getMaxRounds());
|
||||
|
||||
List<AiChatMessageVO> messages = service.listMessages(USER_ID, sessionId);
|
||||
assertEquals(6, messages.size());
|
||||
assertEquals("user", messages.get(0).getRole());
|
||||
assertEquals("assistant", messages.get(1).getRole());
|
||||
assertEquals("问题三:请用一句话总结", messages.get(4).getContent());
|
||||
assertEquals("回复三", messages.get(5).getContent());
|
||||
AiChatSessionMessagesVO payload = service.listMessages(USER_ID, sessionId);
|
||||
assertEquals(6, payload.getMessages().size());
|
||||
assertEquals("user", payload.getMessages().get(0).getRole());
|
||||
assertEquals("assistant", payload.getMessages().get(1).getRole());
|
||||
assertEquals("问题三:请用一句话总结", payload.getMessages().get(4).getContent());
|
||||
assertEquals("回复三", payload.getMessages().get(5).getContent());
|
||||
assertTrue(payload.getModelEnabled());
|
||||
assertNull(payload.getReadOnlyReason());
|
||||
|
||||
PageResult<AiChatSessionVO> page = service.listSessions(USER_ID, 1, 50);
|
||||
assertEquals(1, page.getTotal());
|
||||
assertEquals(3, page.getRecords().get(0).getRoundCount());
|
||||
assertFalse(page.getRecords().get(0).getTitle().isBlank());
|
||||
assertEquals(AiAssistantModelTestSupport.DEFAULT_DISPLAY_NAME, page.getRecords().get(0).getModelDisplayName());
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -104,7 +112,7 @@ class AiChatServiceIntegrationTest {
|
||||
void shouldRejectSixthRound() throws Exception {
|
||||
stubModelReply("r1", "r2", "r3", "r4", "r5");
|
||||
|
||||
Long sessionId = service.createSession(USER_ID).getSessionId();
|
||||
Long sessionId = service.createSession(USER_ID, null).getSessionId();
|
||||
for (int i = 1; i <= 5; i++) {
|
||||
chat(sessionId, "第 " + i + " 轮问题:请简要回答测试内容。");
|
||||
}
|
||||
@@ -122,7 +130,7 @@ class AiChatServiceIntegrationTest {
|
||||
@Tag("p0")
|
||||
@DisplayName("TC-AI-CHAT-IT-03: 模型失败后重试不重复插入 user 消息")
|
||||
void shouldRetryWithoutDuplicateUserMessage() throws Exception {
|
||||
Long sessionId = service.createSession(USER_ID).getSessionId();
|
||||
Long sessionId = service.createSession(USER_ID, null).getSessionId();
|
||||
AiChatSendDTO dto = new AiChatSendDTO();
|
||||
dto.setSessionId(sessionId);
|
||||
dto.setMessage("失败后重试的同一问题内容");
|
||||
@@ -138,7 +146,6 @@ class AiChatServiceIntegrationTest {
|
||||
AiChatResultVO result = service.chat(USER_ID, dto);
|
||||
assertEquals(1, result.getRoundCount());
|
||||
assertEquals(2, store.messagesOf(sessionId).size());
|
||||
// 验证重试未重复插入 user(仅一条 user 消息)
|
||||
long userMsgCount = store.messagesOf(sessionId).stream()
|
||||
.filter(m -> "user".equals(m.getRole()))
|
||||
.count();
|
||||
@@ -149,7 +156,7 @@ class AiChatServiceIntegrationTest {
|
||||
@Tag("p1")
|
||||
@DisplayName("TC-AI-CHAT-IT-04: 用户只能访问自己的会话消息")
|
||||
void shouldIsolateSessionsByUser() {
|
||||
Long sessionId = service.createSession(USER_ID).getSessionId();
|
||||
Long sessionId = service.createSession(USER_ID, null).getSessionId();
|
||||
assertThrows(BusinessException.class, () -> service.listMessages(999L, sessionId));
|
||||
}
|
||||
|
||||
|
||||
+13
-9
@@ -1,12 +1,13 @@
|
||||
package fun.nojava.module.blog.service.impl;
|
||||
|
||||
import fun.nojava.module.blog.config.DashScopeProperties;
|
||||
import fun.nojava.module.blog.dto.AiChatMessageVO;
|
||||
import fun.nojava.module.blog.dto.AiChatResultVO;
|
||||
import fun.nojava.module.blog.dto.AiChatSendDTO;
|
||||
import fun.nojava.module.blog.dto.AiChatSessionMessagesVO;
|
||||
import fun.nojava.module.blog.dto.CreateAiChatSessionVO;
|
||||
import fun.nojava.module.blog.mapper.AiChatMessageMapper;
|
||||
import fun.nojava.module.blog.mapper.AiChatSessionMapper;
|
||||
import fun.nojava.module.blog.service.AiAssistantModelService;
|
||||
import org.junit.jupiter.api.AfterEach;
|
||||
import org.junit.jupiter.api.BeforeEach;
|
||||
import org.junit.jupiter.api.DisplayName;
|
||||
@@ -16,8 +17,6 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
|
||||
import org.springframework.data.redis.core.StringRedisTemplate;
|
||||
import org.springframework.data.redis.core.ValueOperations;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
import static org.mockito.ArgumentMatchers.anyString;
|
||||
import static org.mockito.Mockito.*;
|
||||
@@ -38,6 +37,7 @@ class AiChatServiceLiveIT {
|
||||
private AiSummaryServiceImpl aiSummaryService;
|
||||
private AiChatSessionMapper sessionMapper;
|
||||
private AiChatMessageMapper messageMapper;
|
||||
private AiAssistantModelService aiAssistantModelService;
|
||||
private AiChatServiceImpl service;
|
||||
private final AiChatInMemoryTestSupport store = new AiChatInMemoryTestSupport();
|
||||
|
||||
@@ -61,12 +61,15 @@ class AiChatServiceLiveIT {
|
||||
|
||||
sessionMapper = mock(AiChatSessionMapper.class);
|
||||
messageMapper = mock(AiChatMessageMapper.class);
|
||||
aiAssistantModelService = mock(AiAssistantModelService.class);
|
||||
store.clear();
|
||||
store.wire(sessionMapper, messageMapper);
|
||||
AiAssistantModelTestSupport.wireDefaultModel(aiAssistantModelService);
|
||||
|
||||
aiSummaryService = new AiSummaryServiceImpl(properties, redisTemplate);
|
||||
aiSummaryService.init();
|
||||
service = new AiChatServiceImpl(sessionMapper, messageMapper, aiSummaryService, properties, redisTemplate);
|
||||
service = new AiChatServiceImpl(
|
||||
sessionMapper, messageMapper, aiSummaryService, properties, redisTemplate, aiAssistantModelService);
|
||||
}
|
||||
|
||||
@AfterEach
|
||||
@@ -84,7 +87,7 @@ class AiChatServiceLiveIT {
|
||||
@Tag("live")
|
||||
@DisplayName("TC-AI-CHAT-LIVE-01: 真实调用应完成 2 轮对话并持久化消息")
|
||||
void shouldChatTwoRoundsWithRealModel() {
|
||||
CreateAiChatSessionVO session = service.createSession(FAKE_USER_ID);
|
||||
CreateAiChatSessionVO session = service.createSession(FAKE_USER_ID, null);
|
||||
Long sessionId = session.getSessionId();
|
||||
|
||||
AiChatSendDTO dto1 = new AiChatSendDTO();
|
||||
@@ -104,10 +107,11 @@ class AiChatServiceLiveIT {
|
||||
assertEquals(2, r2.getRoundCount());
|
||||
assertNotNull(r2.getReply());
|
||||
|
||||
List<AiChatMessageVO> messages = service.listMessages(FAKE_USER_ID, sessionId);
|
||||
assertEquals(4, messages.size());
|
||||
assertEquals("user", messages.get(0).getRole());
|
||||
assertEquals("assistant", messages.get(1).getRole());
|
||||
AiChatSessionMessagesVO payload = service.listMessages(FAKE_USER_ID, sessionId);
|
||||
assertEquals(4, payload.getMessages().size());
|
||||
assertEquals("user", payload.getMessages().get(0).getRole());
|
||||
assertEquals("assistant", payload.getMessages().get(1).getRole());
|
||||
assertTrue(payload.getModelEnabled());
|
||||
|
||||
System.out.println("=== AI Chat Live Test ===");
|
||||
System.out.println("sessionId = " + sessionId);
|
||||
|
||||
+10
-33
@@ -5,6 +5,7 @@ import fun.nojava.common.exception.ErrorCode;
|
||||
import fun.nojava.module.blog.config.DashScopeProperties;
|
||||
import fun.nojava.module.blog.dto.GenerateSummaryDTO;
|
||||
import fun.nojava.module.blog.dto.SummaryResultVO;
|
||||
import fun.nojava.module.blog.testutil.TestDataFactory;
|
||||
import org.junit.jupiter.api.BeforeEach;
|
||||
import org.junit.jupiter.api.DisplayName;
|
||||
import org.junit.jupiter.api.Tag;
|
||||
@@ -30,22 +31,14 @@ class AiSummaryServiceImplTest {
|
||||
|
||||
private DashScopeProperties properties;
|
||||
private StringRedisTemplate redisTemplate;
|
||||
private ValueOperations<String, String> valueOps;
|
||||
private TestableAiSummaryService service;
|
||||
|
||||
@SuppressWarnings("unchecked")
|
||||
@BeforeEach
|
||||
void setUp() {
|
||||
properties = new DashScopeProperties();
|
||||
properties.setApiKey("test-key");
|
||||
properties.setBaseUrl("https://example.com/v1");
|
||||
properties.setModel("qwen-turbo");
|
||||
properties.setTimeout(30000);
|
||||
properties.setMaxInput(8000);
|
||||
properties.setRateLimitPerMinute(5);
|
||||
|
||||
properties = TestDataFactory.defaultDashProps();
|
||||
redisTemplate = mock(StringRedisTemplate.class);
|
||||
valueOps = mock(ValueOperations.class);
|
||||
ValueOperations<String, String> valueOps = mock(ValueOperations.class);
|
||||
when(redisTemplate.opsForValue()).thenReturn(valueOps);
|
||||
when(valueOps.increment(anyString())).thenReturn(1L);
|
||||
|
||||
@@ -64,7 +57,6 @@ class AiSummaryServiceImplTest {
|
||||
BusinessException ex = assertThrows(BusinessException.class,
|
||||
() -> service.generateSummary(1L, dto));
|
||||
assertEquals(ErrorCode.AI_CONTENT_TOO_SHORT.getCode(), ex.getCode());
|
||||
verify(valueOps, never()).increment(anyString());
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -78,14 +70,13 @@ class AiSummaryServiceImplTest {
|
||||
BusinessException ex = assertThrows(BusinessException.class,
|
||||
() -> service.generateSummary(1L, dto));
|
||||
assertEquals(ErrorCode.AI_API_KEY_MISSING.getCode(), ex.getCode());
|
||||
verify(valueOps, never()).increment(anyString());
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p0")
|
||||
@DisplayName("TC-AI-SUM-03: 单用户每分钟超过 5 次应触发 RATE_LIMITED")
|
||||
void shouldThrowRateLimitedWhenExceedsLimit() {
|
||||
when(valueOps.increment("ai_summary:rate:7")).thenReturn(6L);
|
||||
when(redisTemplate.opsForValue().increment("ai_summary:rate:7")).thenReturn(6L);
|
||||
|
||||
GenerateSummaryDTO dto = new GenerateSummaryDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文用于触发摘要生成测试用例。");
|
||||
@@ -100,8 +91,8 @@ class AiSummaryServiceImplTest {
|
||||
@Tag("p0")
|
||||
@DisplayName("TC-AI-SUM-04: 首次调用应设置 1 分钟 TTL")
|
||||
void shouldExpireRateLimitKeyOnFirstCall() {
|
||||
when(valueOps.increment("ai_summary:rate:8")).thenReturn(1L);
|
||||
service.stubResponse(buildModelResponse("一段满足长度要求的摘要内容,文字精炼概括核心观点。"));
|
||||
when(redisTemplate.opsForValue().increment("ai_summary:rate:8")).thenReturn(1L);
|
||||
service.stubResponse(TestDataFactory.modelResponse("一段满足长度要求的摘要内容。"));
|
||||
|
||||
GenerateSummaryDTO dto = new GenerateSummaryDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文用于触发摘要生成测试用例。");
|
||||
@@ -115,7 +106,7 @@ class AiSummaryServiceImplTest {
|
||||
@Tag("p0")
|
||||
@DisplayName("TC-AI-SUM-05: 正常返回应填充摘要、模型、tokens、耗时")
|
||||
void shouldReturnSummaryOnSuccess() {
|
||||
service.stubResponse(buildModelResponse("这是一段长度合适的摘要内容,对全文做了精炼概括。"));
|
||||
service.stubResponse(TestDataFactory.modelResponse("这是一段长度合适的摘要内容,对全文做了精炼概括。"));
|
||||
|
||||
GenerateSummaryDTO dto = new GenerateSummaryDTO();
|
||||
dto.setTitle("标题");
|
||||
@@ -135,7 +126,7 @@ class AiSummaryServiceImplTest {
|
||||
@DisplayName("TC-AI-SUM-06: 正文中的 HTML 标签应被剥离后再送入模型")
|
||||
void shouldStripHtmlBeforeCallingModel() {
|
||||
AtomicReference<Map<String, Object>> captured = new AtomicReference<>();
|
||||
service.stubResponseWithCapture(buildModelResponse("摘要内容长度合规并完整。"), captured::set);
|
||||
service.stubResponseWithCapture(TestDataFactory.modelResponse("摘要内容长度合规并完整。"), captured::set);
|
||||
|
||||
GenerateSummaryDTO dto = new GenerateSummaryDTO();
|
||||
dto.setContent("<p>这是 <strong>HTML</strong> 富文本内容,足够长以通过长度校验。</p>");
|
||||
@@ -156,7 +147,7 @@ class AiSummaryServiceImplTest {
|
||||
void shouldTruncateLongContentToMaxInput() {
|
||||
properties.setMaxInput(50);
|
||||
AtomicReference<Map<String, Object>> captured = new AtomicReference<>();
|
||||
service.stubResponseWithCapture(buildModelResponse("摘要内容长度合规并完整。"), captured::set);
|
||||
service.stubResponseWithCapture(TestDataFactory.modelResponse("摘要内容长度合规并完整。"), captured::set);
|
||||
|
||||
GenerateSummaryDTO dto = new GenerateSummaryDTO();
|
||||
dto.setContent("A".repeat(500));
|
||||
@@ -175,7 +166,7 @@ class AiSummaryServiceImplTest {
|
||||
@DisplayName("TC-AI-SUM-08: 模型返回超长内容应截断至 200 字符并加省略号")
|
||||
void shouldTruncateLongModelOutput() {
|
||||
String longSummary = "摘".repeat(300);
|
||||
service.stubResponse(buildModelResponse(longSummary));
|
||||
service.stubResponse(TestDataFactory.modelResponse(longSummary));
|
||||
|
||||
GenerateSummaryDTO dto = new GenerateSummaryDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文用于触发摘要生成测试用例。");
|
||||
@@ -232,20 +223,6 @@ class AiSummaryServiceImplTest {
|
||||
assertEquals(ErrorCode.AI_SERVICE_UNAVAILABLE.getCode(), ex.getCode());
|
||||
}
|
||||
|
||||
private Map<String, Object> buildModelResponse(String summary) {
|
||||
return Map.of(
|
||||
"model", "qwen-turbo",
|
||||
"choices", List.of(
|
||||
Map.of("message", Map.of("role", "assistant", "content", summary))
|
||||
),
|
||||
"usage", Map.of(
|
||||
"prompt_tokens", 123,
|
||||
"completion_tokens", 45,
|
||||
"total_tokens", 168
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* 可观察的 service 子类,劫持 callModel 用于断言传入参数 / 抛出异常。
|
||||
*/
|
||||
|
||||
+39
-85
@@ -2,15 +2,15 @@ package fun.nojava.module.blog.service.impl;
|
||||
|
||||
import fun.nojava.common.exception.BusinessException;
|
||||
import fun.nojava.common.exception.ErrorCode;
|
||||
import fun.nojava.module.blog.config.DashScopeProperties;
|
||||
import fun.nojava.module.blog.dto.RecommendTagsDTO;
|
||||
import fun.nojava.module.blog.dto.TagRecommendResultVO;
|
||||
import fun.nojava.module.blog.testutil.BaseAiServiceTest;
|
||||
import fun.nojava.module.blog.testutil.TestDataFactory;
|
||||
import org.junit.jupiter.api.BeforeEach;
|
||||
import org.junit.jupiter.api.DisplayName;
|
||||
import org.junit.jupiter.api.Tag;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.data.redis.core.StringRedisTemplate;
|
||||
import org.springframework.data.redis.core.ValueOperations;
|
||||
import org.mockito.Mock;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
@@ -19,36 +19,19 @@ import static org.junit.jupiter.api.Assertions.*;
|
||||
import static org.mockito.ArgumentMatchers.*;
|
||||
import static org.mockito.Mockito.*;
|
||||
|
||||
class TagRecommendServiceImplTest {
|
||||
class TagRecommendServiceImplTest extends BaseAiServiceTest {
|
||||
|
||||
private DashScopeProperties properties;
|
||||
private StringRedisTemplate redisTemplate;
|
||||
private ValueOperations<String, String> valueOps;
|
||||
@Mock
|
||||
private AiSummaryServiceImpl mockAiSummaryService;
|
||||
|
||||
private TagRecommendServiceImpl service;
|
||||
|
||||
@SuppressWarnings("unchecked")
|
||||
@BeforeEach
|
||||
void setUp() {
|
||||
properties = new DashScopeProperties();
|
||||
properties.setApiKey("test-key");
|
||||
properties.setBaseUrl("https://example.com/v1");
|
||||
properties.setModel("qwen-turbo");
|
||||
properties.setTimeout(30000);
|
||||
properties.setMaxInput(8000);
|
||||
properties.setRateLimitPerMinute(5);
|
||||
|
||||
redisTemplate = mock(StringRedisTemplate.class);
|
||||
valueOps = mock(ValueOperations.class);
|
||||
when(redisTemplate.opsForValue()).thenReturn(valueOps);
|
||||
when(valueOps.increment(anyString())).thenReturn(1L);
|
||||
|
||||
mockAiSummaryService = mock(AiSummaryServiceImpl.class);
|
||||
|
||||
service = new TagRecommendServiceImpl(mockAiSummaryService, properties, redisTemplate);
|
||||
service = new TagRecommendServiceImpl(mockAiSummaryService, dashProps, redisTemplate);
|
||||
}
|
||||
|
||||
// ==================== P0 测试 ====================
|
||||
// ==================== P0 ====================
|
||||
|
||||
@Test
|
||||
@Tag("p0")
|
||||
@@ -68,10 +51,10 @@ class TagRecommendServiceImplTest {
|
||||
@Tag("p0")
|
||||
@DisplayName("TC-TAG-REC-02: 未配置 API Key 应抛出 AI_API_KEY_MISSING")
|
||||
void shouldRejectWhenApiKeyMissing() {
|
||||
properties.setApiKey("");
|
||||
dashProps.setApiKey("");
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class,
|
||||
() -> service.recommendTags(1L, dto));
|
||||
@@ -86,7 +69,7 @@ class TagRecommendServiceImplTest {
|
||||
when(valueOps.increment("tag_recommend:rate:7")).thenReturn(6L);
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class,
|
||||
() -> service.recommendTags(7L, dto));
|
||||
@@ -99,11 +82,11 @@ class TagRecommendServiceImplTest {
|
||||
@DisplayName("TC-TAG-REC-04: 正常返回应填充 tags、model、tokens、耗时")
|
||||
void shouldReturnTagsOnSuccess() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("Java, Spring Boot, 微服务"));
|
||||
.thenReturn(TestDataFactory.modelResponse("Java, Spring Boot, 微服务"));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setTitle("Spring Cloud 微服务实践");
|
||||
dto.setContent("这是一篇关于微服务架构的详细技术文章,涵盖服务注册与发现、配置中心等内容。");
|
||||
dto.setContent("关于微服务架构的详细技术文章,涵盖服务注册与发现、配置中心等内容。");
|
||||
|
||||
TagRecommendResultVO result = service.recommendTags(1L, dto);
|
||||
|
||||
@@ -123,13 +106,12 @@ class TagRecommendServiceImplTest {
|
||||
@DisplayName("TC-TAG-REC-05: 标签数量不应超过 MAX_TAGS(5个)")
|
||||
void shouldLimitTagsToMax() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("Java, Spring, 微服务, Docker, K8s, 多余标签, 更多"));
|
||||
.thenReturn(TestDataFactory.modelResponse("Java, Spring, 微服务, Docker, K8s, 多余, 更多"));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
TagRecommendResultVO result = service.recommendTags(1L, dto);
|
||||
|
||||
assertEquals(5, result.getTags().size());
|
||||
}
|
||||
|
||||
@@ -138,17 +120,15 @@ class TagRecommendServiceImplTest {
|
||||
@DisplayName("TC-TAG-REC-06: 支持中英文逗号混合分隔")
|
||||
void shouldParseCommaAndChineseComma() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("Java,Spring Boot,微服务"));
|
||||
.thenReturn(TestDataFactory.modelResponse("Java,Spring Boot,微服务"));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
TagRecommendResultVO result = service.recommendTags(1L, dto);
|
||||
|
||||
assertEquals(3, result.getTags().size());
|
||||
assertEquals("Java", result.getTags().get(0));
|
||||
assertEquals("Spring Boot", result.getTags().get(1));
|
||||
assertEquals("微服务", result.getTags().get(2));
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -156,13 +136,12 @@ class TagRecommendServiceImplTest {
|
||||
@DisplayName("TC-TAG-REC-07: 空白标签应被过滤")
|
||||
void shouldFilterEmptyTags() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("Java, , Spring, , 微服务"));
|
||||
.thenReturn(TestDataFactory.modelResponse("Java, , Spring, , 微服务"));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
TagRecommendResultVO result = service.recommendTags(1L, dto);
|
||||
|
||||
assertEquals(3, result.getTags().size());
|
||||
}
|
||||
|
||||
@@ -171,7 +150,7 @@ class TagRecommendServiceImplTest {
|
||||
@DisplayName("TC-TAG-REC-08: 正文中的 HTML 标签应被剥离")
|
||||
void shouldStripHtmlBeforeCallingModel() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("Java, Spring Boot"));
|
||||
.thenReturn(TestDataFactory.modelResponse("Java, Spring Boot"));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("<p>这是一段 <strong>HTML</strong> 富文本,足够长以通过校验。</p>");
|
||||
@@ -184,7 +163,6 @@ class TagRecommendServiceImplTest {
|
||||
String userContent = (String) messages.get(1).get("content");
|
||||
assertFalse(userContent.contains("<p>"), "HTML 标签未被剥离: " + userContent);
|
||||
assertFalse(userContent.contains("<strong>"), "HTML 标签未被剥离: " + userContent);
|
||||
assertTrue(userContent.contains("HTML"), "纯文本内容缺失");
|
||||
return true;
|
||||
}));
|
||||
}
|
||||
@@ -199,7 +177,7 @@ class TagRecommendServiceImplTest {
|
||||
when(mockAiSummaryService.callModel(any())).thenReturn(response);
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class,
|
||||
() -> service.recommendTags(1L, dto));
|
||||
@@ -214,28 +192,27 @@ class TagRecommendServiceImplTest {
|
||||
.thenThrow(new RuntimeException("网络异常"));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class,
|
||||
() -> service.recommendTags(1L, dto));
|
||||
assertEquals(ErrorCode.AI_SERVICE_UNAVAILABLE.getCode(), ex.getCode());
|
||||
}
|
||||
|
||||
// ==================== P1 测试 ====================
|
||||
// ==================== P1 ====================
|
||||
|
||||
@Test
|
||||
@Tag("p1")
|
||||
@DisplayName("TC-TAG-REC-11: 标题为 null 时应正常调用模型")
|
||||
void shouldWorkWhenTitleIsNull() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("标签A, 标签B"));
|
||||
.thenReturn(TestDataFactory.modelResponse("标签A, 标签B"));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setTitle(null);
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
TagRecommendResultVO result = service.recommendTags(1L, dto);
|
||||
|
||||
assertNotNull(result.getTags());
|
||||
assertEquals(2, result.getTags().size());
|
||||
}
|
||||
@@ -244,9 +221,9 @@ class TagRecommendServiceImplTest {
|
||||
@Tag("p1")
|
||||
@DisplayName("TC-TAG-REC-12: 长正文应截断到 maxInput 上限")
|
||||
void shouldTruncateLongContentToMaxInput() {
|
||||
properties.setMaxInput(50);
|
||||
dashProps.setMaxInput(50);
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("Java, Spring"));
|
||||
.thenReturn(TestDataFactory.modelResponse("Java, Spring"));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("X".repeat(500));
|
||||
@@ -257,9 +234,7 @@ class TagRecommendServiceImplTest {
|
||||
@SuppressWarnings("unchecked")
|
||||
List<Map<String, Object>> messages = (List<Map<String, Object>>) requestBody.get("messages");
|
||||
String userContent = (String) messages.get(1).get("content");
|
||||
// 正文部分应该被截断到 50 个字符
|
||||
long xCount = userContent.chars().filter(c -> c == 'X').count();
|
||||
assertEquals(50, xCount, "正文应截断到 maxInput 字符");
|
||||
assertEquals(50, userContent.chars().filter(c -> c == 'X').count());
|
||||
return true;
|
||||
}));
|
||||
}
|
||||
@@ -270,10 +245,10 @@ class TagRecommendServiceImplTest {
|
||||
void shouldExpireRateLimitKeyOnFirstCall() {
|
||||
when(valueOps.increment("tag_recommend:rate:8")).thenReturn(1L);
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("Java, Spring"));
|
||||
.thenReturn(TestDataFactory.modelResponse("Java, Spring"));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
service.recommendTags(8L, dto);
|
||||
|
||||
@@ -285,16 +260,14 @@ class TagRecommendServiceImplTest {
|
||||
@DisplayName("TC-TAG-REC-14: userId 为 null 时不应限流,应正常调用模型")
|
||||
void shouldNotRateLimitWhenUserIdNull() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("Java, Spring"));
|
||||
.thenReturn(TestDataFactory.modelResponse("Java, Spring"));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
TagRecommendResultVO result = service.recommendTags(null, dto);
|
||||
|
||||
assertNotNull(result.getTags());
|
||||
assertEquals(2, result.getTags().size());
|
||||
verify(mockAiSummaryService).callModel(any());
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -302,10 +275,10 @@ class TagRecommendServiceImplTest {
|
||||
@DisplayName("TC-TAG-REC-15: 模型返回 content 为 blank 应映射为 AI_SERVICE_UNAVAILABLE")
|
||||
void shouldFailWhenContentIsBlank() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse(" "));
|
||||
.thenReturn(TestDataFactory.modelResponse(" "));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
BusinessException ex = assertThrows(BusinessException.class,
|
||||
() -> service.recommendTags(1L, dto));
|
||||
@@ -325,10 +298,9 @@ class TagRecommendServiceImplTest {
|
||||
when(mockAiSummaryService.callModel(any())).thenReturn(response);
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
TagRecommendResultVO result = service.recommendTags(1L, dto);
|
||||
|
||||
assertNull(result.getPromptTokens());
|
||||
assertNull(result.getCompletionTokens());
|
||||
}
|
||||
@@ -338,13 +310,12 @@ class TagRecommendServiceImplTest {
|
||||
@DisplayName("TC-TAG-REC-17: 模型返回单个标签应正常解析")
|
||||
void shouldWorkWithSingleTag() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("全栈开发"));
|
||||
.thenReturn(TestDataFactory.modelResponse("全栈开发"));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
TagRecommendResultVO result = service.recommendTags(1L, dto);
|
||||
|
||||
assertEquals(1, result.getTags().size());
|
||||
assertEquals("全栈开发", result.getTags().get(0));
|
||||
}
|
||||
@@ -354,31 +325,14 @@ class TagRecommendServiceImplTest {
|
||||
@DisplayName("TC-TAG-REC-18: 模型返回前后带空格的标签应去除空格")
|
||||
void shouldTrimTagWhitespace() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse(" Java , Spring Boot , 微服务 "));
|
||||
.thenReturn(TestDataFactory.modelResponse(" Java , Spring Boot , 微服务 "));
|
||||
|
||||
RecommendTagsDTO dto = new RecommendTagsDTO();
|
||||
dto.setContent("这是一篇足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
dto.setContent("足够长的文章正文内容,用于触发标签推荐测试用例。");
|
||||
|
||||
TagRecommendResultVO result = service.recommendTags(1L, dto);
|
||||
|
||||
assertEquals("Java", result.getTags().get(0));
|
||||
assertEquals("Spring Boot", result.getTags().get(1));
|
||||
assertEquals("微服务", result.getTags().get(2));
|
||||
}
|
||||
|
||||
// ==================== 辅助方法 ====================
|
||||
|
||||
private Map<String, Object> buildModelResponse(String tagsContent) {
|
||||
return Map.of(
|
||||
"model", "qwen-turbo",
|
||||
"choices", List.of(
|
||||
Map.of("message", Map.of("role", "assistant", "content", tagsContent))
|
||||
),
|
||||
"usage", Map.of(
|
||||
"prompt_tokens", 123,
|
||||
"completion_tokens", 45,
|
||||
"total_tokens", 168
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+43
-98
@@ -2,16 +2,16 @@ package fun.nojava.module.blog.service.impl;
|
||||
|
||||
import fun.nojava.common.exception.BusinessException;
|
||||
import fun.nojava.common.exception.ErrorCode;
|
||||
import fun.nojava.module.blog.config.DashScopeProperties;
|
||||
import fun.nojava.module.blog.dto.DiffSegmentVO;
|
||||
import fun.nojava.module.blog.dto.PolishResultVO;
|
||||
import fun.nojava.module.blog.dto.PolishTextDTO;
|
||||
import fun.nojava.module.blog.testutil.BaseAiServiceTest;
|
||||
import fun.nojava.module.blog.testutil.TestDataFactory;
|
||||
import org.junit.jupiter.api.BeforeEach;
|
||||
import org.junit.jupiter.api.DisplayName;
|
||||
import org.junit.jupiter.api.Tag;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.data.redis.core.StringRedisTemplate;
|
||||
import org.springframework.data.redis.core.ValueOperations;
|
||||
import org.mockito.Mock;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
@@ -20,36 +20,19 @@ import static org.junit.jupiter.api.Assertions.*;
|
||||
import static org.mockito.ArgumentMatchers.*;
|
||||
import static org.mockito.Mockito.*;
|
||||
|
||||
class TextPolishServiceImplTest {
|
||||
class TextPolishServiceImplTest extends BaseAiServiceTest {
|
||||
|
||||
private DashScopeProperties properties;
|
||||
private StringRedisTemplate redisTemplate;
|
||||
private ValueOperations<String, String> valueOps;
|
||||
@Mock
|
||||
private AiSummaryServiceImpl mockAiSummaryService;
|
||||
|
||||
private TextPolishServiceImpl service;
|
||||
|
||||
@SuppressWarnings("unchecked")
|
||||
@BeforeEach
|
||||
void setUp() {
|
||||
properties = new DashScopeProperties();
|
||||
properties.setApiKey("test-key");
|
||||
properties.setBaseUrl("https://example.com/v1");
|
||||
properties.setModel("qwen-turbo");
|
||||
properties.setTimeout(30000);
|
||||
properties.setMaxInput(6000);
|
||||
properties.setRateLimitPerMinute(5);
|
||||
|
||||
redisTemplate = mock(StringRedisTemplate.class);
|
||||
valueOps = mock(ValueOperations.class);
|
||||
when(redisTemplate.opsForValue()).thenReturn(valueOps);
|
||||
when(valueOps.increment(anyString())).thenReturn(1L);
|
||||
|
||||
mockAiSummaryService = mock(AiSummaryServiceImpl.class);
|
||||
|
||||
service = new TextPolishServiceImpl(mockAiSummaryService, properties, redisTemplate);
|
||||
service = new TextPolishServiceImpl(mockAiSummaryService, dashProps, redisTemplate);
|
||||
}
|
||||
|
||||
// ==================== P0 测试 ====================
|
||||
// ==================== P0 ====================
|
||||
|
||||
@Test
|
||||
@Tag("p0")
|
||||
@@ -68,7 +51,7 @@ class TextPolishServiceImplTest {
|
||||
@Tag("p0")
|
||||
@DisplayName("TC-POLISH-02: 未配置 API Key 应抛出 AI_API_KEY_MISSING")
|
||||
void shouldRejectWhenApiKeyMissing() {
|
||||
properties.setApiKey("");
|
||||
dashProps.setApiKey("");
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText("这是一段足够长的文本内容,用于触发文本润色测试用例验证。");
|
||||
@@ -101,7 +84,7 @@ class TextPolishServiceImplTest {
|
||||
String originalText = "这是一篇关于微服务架构的详细技术文章,涵盖服务注册与发现、配置中心等内容。";
|
||||
String polishedText = "这是一篇关于微服务架构的详细技术文章,覆盖了服务注册与发现、配置中心等内容。";
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse(polishedText));
|
||||
.thenReturn(TestDataFactory.modelResponse(polishedText));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText(originalText);
|
||||
@@ -114,8 +97,6 @@ class TextPolishServiceImplTest {
|
||||
assertNotNull(result.getDiffSegments());
|
||||
assertFalse(result.getDiffSegments().isEmpty());
|
||||
assertEquals(0, result.getDiffSegments().get(0).getIndex());
|
||||
assertEquals(originalText, result.getDiffSegments().get(0).getOriginal());
|
||||
assertEquals(polishedText, result.getDiffSegments().get(0).getPolished());
|
||||
assertEquals("qwen-turbo", result.getModel());
|
||||
assertEquals(123, result.getPromptTokens());
|
||||
assertEquals(45, result.getCompletionTokens());
|
||||
@@ -129,7 +110,7 @@ class TextPolishServiceImplTest {
|
||||
String originalText = "第一段内容。\n\n第二段内容。\n\n第三段内容。";
|
||||
String polishedText = "第一段润色后。\n\n第二段润色后。\n\n第三段润色后。";
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse(polishedText));
|
||||
.thenReturn(TestDataFactory.modelResponse(polishedText));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText(originalText);
|
||||
@@ -139,10 +120,6 @@ class TextPolishServiceImplTest {
|
||||
assertEquals(3, result.getDiffSegments().size());
|
||||
assertEquals("第一段内容。", result.getDiffSegments().get(0).getOriginal());
|
||||
assertEquals("第一段润色后。", result.getDiffSegments().get(0).getPolished());
|
||||
assertEquals("第二段内容。", result.getDiffSegments().get(1).getOriginal());
|
||||
assertEquals("第二段润色后。", result.getDiffSegments().get(1).getPolished());
|
||||
assertEquals("第三段内容。", result.getDiffSegments().get(2).getOriginal());
|
||||
assertEquals("第三段润色后。", result.getDiffSegments().get(2).getPolished());
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -181,13 +158,12 @@ class TextPolishServiceImplTest {
|
||||
@Tag("p0")
|
||||
@DisplayName("TC-POLISH-08: 长文本应截断到 maxInput 上限")
|
||||
void shouldTruncateLongTextToMaxInput() {
|
||||
properties.setMaxInput(50);
|
||||
dashProps.setMaxInput(50);
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("润色后的文本内容。"));
|
||||
String longText = "X".repeat(500);
|
||||
.thenReturn(TestDataFactory.modelResponse("润色后的文本内容。"));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText(longText);
|
||||
dto.setText("X".repeat(500));
|
||||
|
||||
service.polishText(1L, dto);
|
||||
|
||||
@@ -195,8 +171,7 @@ class TextPolishServiceImplTest {
|
||||
@SuppressWarnings("unchecked")
|
||||
List<Map<String, Object>> messages = (List<Map<String, Object>>) requestBody.get("messages");
|
||||
String userContent = (String) messages.get(1).get("content");
|
||||
long xCount = userContent.chars().filter(c -> c == 'X').count();
|
||||
assertEquals(50, xCount, "文本应截断到 maxInput 字符");
|
||||
assertEquals(50, userContent.chars().filter(c -> c == 'X').count());
|
||||
return true;
|
||||
}));
|
||||
}
|
||||
@@ -207,7 +182,7 @@ class TextPolishServiceImplTest {
|
||||
void shouldExpireRateLimitKeyOnFirstCall() {
|
||||
when(valueOps.increment("text_polish:rate:8")).thenReturn(1L);
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("润色后的文本内容。"));
|
||||
.thenReturn(TestDataFactory.modelResponse("润色后的文本内容。"));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText("这是一段足够长的文本内容,用于触发文本润色测试用例验证。");
|
||||
@@ -222,25 +197,23 @@ class TextPolishServiceImplTest {
|
||||
@DisplayName("TC-POLISH-10: userId 为 null 时不应限流,应正常调用模型")
|
||||
void shouldNotRateLimitWhenUserIdNull() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("润色后的文本内容。"));
|
||||
.thenReturn(TestDataFactory.modelResponse("润色后的文本内容。"));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText("这是一段足够长的文本内容,用于触发文本润色测试用例验证。");
|
||||
|
||||
PolishResultVO result = service.polishText(null, dto);
|
||||
|
||||
assertNotNull(result.getPolished());
|
||||
verify(mockAiSummaryService).callModel(any());
|
||||
}
|
||||
|
||||
// ==================== P1 测试 ====================
|
||||
// ==================== P1 ====================
|
||||
|
||||
@Test
|
||||
@Tag("p1")
|
||||
@DisplayName("TC-POLISH-11: 模型返回 content 为 blank 应映射为 AI_SERVICE_UNAVAILABLE")
|
||||
void shouldFailWhenContentIsBlank() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse(" "));
|
||||
.thenReturn(TestDataFactory.modelResponse(" "));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText("这是一段足够长的文本内容,用于触发文本润色测试用例验证。");
|
||||
@@ -286,7 +259,6 @@ class TextPolishServiceImplTest {
|
||||
dto.setText("这是一段足够长的文本内容,用于触发文本润色测试用例验证。");
|
||||
|
||||
PolishResultVO result = service.polishText(1L, dto);
|
||||
|
||||
assertNull(result.getPromptTokens());
|
||||
assertNull(result.getCompletionTokens());
|
||||
}
|
||||
@@ -295,35 +267,31 @@ class TextPolishServiceImplTest {
|
||||
@Tag("p1")
|
||||
@DisplayName("TC-POLISH-14: 段落数差异超过 2 倍应降级为全文单段对比(原文段落多)")
|
||||
void shouldFallbackWhenParagraphCountDiffersTooMuchMoreOriginal() {
|
||||
String originalText = "第一段落内容足够长。\n\n第二段落内容足够长。\n\n第三段落内容足够长。\n\n第四段落内容足够长。\n\n第五段落内容足够长。";
|
||||
String polishedText = "全部合并为一段润色后的文本内容。";
|
||||
String originalText = "第一段。\n\n第二段。\n\n第三段。\n\n第四段。\n\n第五段。";
|
||||
String polishedText = "全部合并为一段润色后的文本。";
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse(polishedText));
|
||||
.thenReturn(TestDataFactory.modelResponse(polishedText));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText(originalText);
|
||||
|
||||
PolishResultVO result = service.polishText(1L, dto);
|
||||
|
||||
assertEquals(1, result.getDiffSegments().size(), "应降级为单段对比");
|
||||
assertEquals(originalText, result.getDiffSegments().get(0).getOriginal());
|
||||
assertEquals(polishedText, result.getDiffSegments().get(0).getPolished());
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p1")
|
||||
@DisplayName("TC-POLISH-15: 段落数差异超过 2 倍应降级为全文单段对比(润色段落多)")
|
||||
void shouldFallbackWhenParagraphCountDiffersTooMuchMorePolished() {
|
||||
String originalText = "只有一段但足够长的文本内容用于润色处理。";
|
||||
String polishedText = "润色后第一段文本内容足够长。\n\n润色后第二段文本内容足够长。\n\n润色后第三段文本内容足够长。\n\n润色后第四段文本内容足够长。";
|
||||
String originalText = "只有一段但足够长的文本内容用于测试验证。";
|
||||
String polishedText = "一段。\n\n二段。\n\n三段。\n\n四段。";
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse(polishedText));
|
||||
.thenReturn(TestDataFactory.modelResponse(polishedText));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText(originalText);
|
||||
|
||||
PolishResultVO result = service.polishText(1L, dto);
|
||||
|
||||
assertEquals(1, result.getDiffSegments().size(), "应降级为单段对比");
|
||||
}
|
||||
|
||||
@@ -331,27 +299,26 @@ class TextPolishServiceImplTest {
|
||||
@Tag("p1")
|
||||
@DisplayName("TC-POLISH-16: 段落数略有差异(不超过2倍)时按较少段落数对齐")
|
||||
void shouldAlignToMinParagraphsWhenSlightDifference() {
|
||||
String originalText = "第一篇原文内容足够长。\n\n第二篇原文内容足够长。";
|
||||
String polishedText = "第一篇润色后内容足够长。\n\n第二篇润色后内容足够长。\n\n多余的第三段落内容。";
|
||||
String originalText = "第一篇原文内容足够长。\n\n第二篇原文内容也足够长。";
|
||||
String polishedText = "第一篇润色后。\n\n第二篇润色后。\n\n多余段落。";
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse(polishedText));
|
||||
.thenReturn(TestDataFactory.modelResponse(polishedText));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText(originalText);
|
||||
|
||||
PolishResultVO result = service.polishText(1L, dto);
|
||||
|
||||
assertEquals(2, result.getDiffSegments().size(), "应按较少段落数(2)对齐");
|
||||
assertEquals(2, result.getDiffSegments().size());
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("p1")
|
||||
@DisplayName("TC-POLISH-17: 单段文本正常返回")
|
||||
void shouldWorkWithSingleParagraph() {
|
||||
String originalText = "这是一段需要润色的单段文本内容,存在一些语法和表达的问题。";
|
||||
String polishedText = "这是一段已润色的单段文本内容,修正了语法和表达的问题。";
|
||||
String originalText = "这是一段需要润色的单段文本内容,存在一些语法问题。";
|
||||
String polishedText = "这是一段已润色的单段文本内容,修正了语法问题。";
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse(polishedText));
|
||||
.thenReturn(TestDataFactory.modelResponse(polishedText));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText(originalText);
|
||||
@@ -367,18 +334,16 @@ class TextPolishServiceImplTest {
|
||||
@Tag("p1")
|
||||
@DisplayName("TC-POLISH-18: 包含空段落的文本应过滤空段")
|
||||
void shouldFilterEmptyParagraphs() {
|
||||
String originalText = "第一段落内容足够长。\n\n\n\n第二段落内容足够长。";
|
||||
String polishedText = "第一段润色后内容。\n\n第二段润色后内容。";
|
||||
String originalText = "第一段落内容足够长。\n\n\n\n第二段落内容也足够长。";
|
||||
String polishedText = "第一段润色后。\n\n第二段润色后。";
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse(polishedText));
|
||||
.thenReturn(TestDataFactory.modelResponse(polishedText));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText(originalText);
|
||||
|
||||
PolishResultVO result = service.polishText(1L, dto);
|
||||
|
||||
assertFalse(result.getDiffSegments().isEmpty());
|
||||
// 过滤后应有2段有效内容
|
||||
assertTrue(result.getDiffSegments().size() >= 1);
|
||||
}
|
||||
|
||||
@@ -386,17 +351,16 @@ class TextPolishServiceImplTest {
|
||||
@Tag("p1")
|
||||
@DisplayName("TC-POLISH-19: 文本前后空格应被 trim")
|
||||
void shouldTrimTextBeforeProcessing() {
|
||||
String originalText = " 这是一段前后有空格的文本,足够长以通过长度校验测试。 ";
|
||||
String originalText = " 前后有空格的文本,足够长以通过长度校验。 ";
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("润色后的内容。"));
|
||||
.thenReturn(TestDataFactory.modelResponse("润色后的内容。"));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText(originalText);
|
||||
|
||||
PolishResultVO result = service.polishText(1L, dto);
|
||||
|
||||
assertFalse(result.getOriginal().startsWith(" "), "开头空格应被 trim");
|
||||
assertFalse(result.getOriginal().endsWith(" "), "结尾空格应被 trim");
|
||||
assertFalse(result.getOriginal().startsWith(" "));
|
||||
assertFalse(result.getOriginal().endsWith(" "));
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -437,7 +401,7 @@ class TextPolishServiceImplTest {
|
||||
@DisplayName("TC-POLISH-22: 传递的 model、temperature、messages 结构应正确")
|
||||
void shouldBuildCorrectModelRequestBody() {
|
||||
when(mockAiSummaryService.callModel(any()))
|
||||
.thenReturn(buildModelResponse("润色后的文本。"));
|
||||
.thenReturn(TestDataFactory.modelResponse("润色后的文本。"));
|
||||
|
||||
PolishTextDTO dto = new PolishTextDTO();
|
||||
dto.setText("这是一段足够长的文本内容,用于触发文本润色测试用例验证。");
|
||||
@@ -451,15 +415,14 @@ class TextPolishServiceImplTest {
|
||||
List<Map<String, Object>> messages = (List<Map<String, Object>>) requestBody.get("messages");
|
||||
assertEquals(2, messages.size());
|
||||
assertEquals("system", messages.get(0).get("role"));
|
||||
assertEquals("user", messages.get(1).get("role"));
|
||||
String userContent = (String) messages.get(1).get("content");
|
||||
assertTrue(userContent.contains("请润色下面这段中文文本"), "User prompt 应包含润色指令");
|
||||
assertTrue(userContent.contains("测试用例验证"), "User prompt 应包含正文内容");
|
||||
assertTrue(userContent.contains("请润色下面这段中文文本"));
|
||||
assertTrue(userContent.contains("测试用例验证"));
|
||||
return true;
|
||||
}));
|
||||
}
|
||||
|
||||
// ==================== 段落对齐专项测试 ====================
|
||||
// ==================== 段落对齐专项 ====================
|
||||
|
||||
@Test
|
||||
@Tag("p1")
|
||||
@@ -475,8 +438,6 @@ class TextPolishServiceImplTest {
|
||||
void alignParagraphs_identicalText() {
|
||||
List<DiffSegmentVO> result = service.alignParagraphs("相同内容", "相同内容");
|
||||
assertEquals(1, result.size());
|
||||
assertEquals("相同内容", result.get(0).getOriginal());
|
||||
assertEquals("相同内容", result.get(0).getPolished());
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -486,20 +447,4 @@ class TextPolishServiceImplTest {
|
||||
List<DiffSegmentVO> result = service.alignParagraphs("", "");
|
||||
assertEquals(1, result.size());
|
||||
}
|
||||
|
||||
// ==================== 辅助方法 ====================
|
||||
|
||||
private Map<String, Object> buildModelResponse(String polishedText) {
|
||||
return Map.of(
|
||||
"model", "qwen-turbo",
|
||||
"choices", List.of(
|
||||
Map.of("message", Map.of("role", "assistant", "content", polishedText))
|
||||
),
|
||||
"usage", Map.of(
|
||||
"prompt_tokens", 123,
|
||||
"completion_tokens", 45,
|
||||
"total_tokens", 168
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
package fun.nojava.module.blog.testutil;
|
||||
|
||||
import fun.nojava.module.blog.config.DashScopeProperties;
|
||||
import org.junit.jupiter.api.BeforeEach;
|
||||
import org.junit.jupiter.api.extension.ExtendWith;
|
||||
import org.mockito.Mock;
|
||||
import org.mockito.junit.jupiter.MockitoExtension;
|
||||
import org.springframework.data.redis.core.StringRedisTemplate;
|
||||
import org.springframework.data.redis.core.ValueOperations;
|
||||
|
||||
import java.util.Map;
|
||||
|
||||
import static org.mockito.ArgumentMatchers.anyString;
|
||||
import static org.mockito.Mockito.lenient;
|
||||
import static org.mockito.Mockito.when;
|
||||
|
||||
/**
|
||||
* 为需要 {@link DashScopeProperties} 和 Redis 限流的 AI 服务测试提供共享 setup。
|
||||
* 子类通过 {@code @ExtendWith(MockitoExtension.class)} 获得 Mockito 支持(本类声明后子类自动继承)。
|
||||
*/
|
||||
@ExtendWith(MockitoExtension.class)
|
||||
public abstract class BaseAiServiceTest {
|
||||
|
||||
protected DashScopeProperties dashProps;
|
||||
|
||||
@Mock
|
||||
protected StringRedisTemplate redisTemplate;
|
||||
|
||||
@Mock
|
||||
protected ValueOperations<String, String> valueOps;
|
||||
|
||||
@BeforeEach
|
||||
void setUpBase() {
|
||||
dashProps = TestDataFactory.defaultDashProps();
|
||||
lenient().when(redisTemplate.opsForValue()).thenReturn(valueOps);
|
||||
lenient().when(valueOps.increment(anyString())).thenReturn(1L);
|
||||
}
|
||||
|
||||
/** 构建 DashScope 成功响应,默认模型 qwen-turbo。 */
|
||||
protected static Map<String, Object> modelReply(String content) {
|
||||
return TestDataFactory.modelResponse(content);
|
||||
}
|
||||
|
||||
/** 构建 DashScope 成功响应,指定 modelCode。 */
|
||||
protected static Map<String, Object> modelReply(String content, String modelCode) {
|
||||
return TestDataFactory.modelResponse(content, modelCode);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,129 @@
|
||||
package fun.nojava.module.blog.testutil;
|
||||
|
||||
import fun.nojava.module.blog.config.DashScopeProperties;
|
||||
import fun.nojava.module.blog.dto.AiAssistantModelSaveDTO;
|
||||
import fun.nojava.module.blog.dto.AiChatSendDTO;
|
||||
import fun.nojava.module.blog.entity.AiAssistantModel;
|
||||
import fun.nojava.module.blog.entity.AiChatMessage;
|
||||
import fun.nojava.module.blog.entity.AiChatSession;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
public final class TestDataFactory {
|
||||
|
||||
private TestDataFactory() {}
|
||||
|
||||
// === DashScopeProperties ===
|
||||
|
||||
public static DashScopeProperties defaultDashProps() {
|
||||
DashScopeProperties props = new DashScopeProperties();
|
||||
props.setApiKey("test-key");
|
||||
props.setBaseUrl("https://example.com/v1");
|
||||
props.setModel("qwen-turbo");
|
||||
props.setTimeout(30000);
|
||||
props.setMaxInput(8000);
|
||||
props.setRateLimitPerMinute(5);
|
||||
return props;
|
||||
}
|
||||
|
||||
// === Model response (DashScope OpenAI-compatible) ===
|
||||
|
||||
public static Map<String, Object> modelResponse(String content) {
|
||||
return modelResponse(content, "qwen-turbo");
|
||||
}
|
||||
|
||||
public static Map<String, Object> modelResponse(String content, String modelCode) {
|
||||
return Map.of(
|
||||
"model", modelCode,
|
||||
"choices", List.of(Map.of("message", Map.of("role", "assistant", "content", content))),
|
||||
"usage", Map.of("prompt_tokens", 123, "completion_tokens", 45, "total_tokens", 168)
|
||||
);
|
||||
}
|
||||
|
||||
public static Map<String, Object> modelResponse(String content, int promptTokens, int completionTokens) {
|
||||
return Map.of(
|
||||
"model", "qwen-turbo",
|
||||
"choices", List.of(Map.of("message", Map.of("role", "assistant", "content", content))),
|
||||
"usage", Map.of("prompt_tokens", promptTokens, "completion_tokens", completionTokens)
|
||||
);
|
||||
}
|
||||
|
||||
// === AiAssistantModel ===
|
||||
|
||||
public static AiAssistantModel defaultModel() {
|
||||
AiAssistantModel m = new AiAssistantModel();
|
||||
m.setId(1L);
|
||||
m.setModelCode("qwen-turbo");
|
||||
m.setDisplayName("默认模型");
|
||||
m.setEnabled(true);
|
||||
m.setIsDefault(true);
|
||||
m.setDeleted(0);
|
||||
return m;
|
||||
}
|
||||
|
||||
public static AiAssistantModel aModel(long id, String code, String displayName) {
|
||||
AiAssistantModel m = new AiAssistantModel();
|
||||
m.setId(id);
|
||||
m.setModelCode(code);
|
||||
m.setDisplayName(displayName);
|
||||
m.setEnabled(true);
|
||||
m.setDeleted(0);
|
||||
return m;
|
||||
}
|
||||
|
||||
// === AiChatSession ===
|
||||
|
||||
public static AiChatSession aSession(long id, long userId, long modelId) {
|
||||
AiChatSession s = new AiChatSession();
|
||||
s.setId(id);
|
||||
s.setUserId(userId);
|
||||
s.setModelId(modelId);
|
||||
s.setTitle("测试会话");
|
||||
s.setRoundCount(0);
|
||||
s.setDeleted(0);
|
||||
return s;
|
||||
}
|
||||
|
||||
public static AiChatSession aSession(long id, long userId, long modelId, int roundCount) {
|
||||
AiChatSession s = aSession(id, userId, modelId);
|
||||
s.setRoundCount(roundCount);
|
||||
return s;
|
||||
}
|
||||
|
||||
// === AiChatMessage ===
|
||||
|
||||
public static AiChatMessage userMessage(long sessionId, String content, int sortOrder) {
|
||||
AiChatMessage m = new AiChatMessage();
|
||||
m.setSessionId(sessionId);
|
||||
m.setRole("user");
|
||||
m.setContent(content);
|
||||
m.setSortOrder(sortOrder);
|
||||
return m;
|
||||
}
|
||||
|
||||
public static AiChatMessage assistantMessage(long sessionId, String content, int sortOrder) {
|
||||
AiChatMessage m = new AiChatMessage();
|
||||
m.setSessionId(sessionId);
|
||||
m.setRole("assistant");
|
||||
m.setContent(content);
|
||||
m.setSortOrder(sortOrder);
|
||||
return m;
|
||||
}
|
||||
|
||||
// === DTOs ===
|
||||
|
||||
public static AiChatSendDTO aChatSendDTO(long sessionId, String message) {
|
||||
AiChatSendDTO dto = new AiChatSendDTO();
|
||||
dto.setSessionId(sessionId);
|
||||
dto.setMessage(message);
|
||||
return dto;
|
||||
}
|
||||
|
||||
public static AiAssistantModelSaveDTO aModelSaveDTO(String code, String displayName) {
|
||||
AiAssistantModelSaveDTO dto = new AiAssistantModelSaveDTO();
|
||||
dto.setModelCode(code);
|
||||
dto.setDisplayName(displayName);
|
||||
return dto;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
# JUnit Platform 配置
|
||||
# 放在 src/test/resources/ 下自动生效
|
||||
|
||||
# 测试发现:支持 *Test, *Tests, *IT 后缀
|
||||
junit.platform.discovery.includeClassNamePattern=^.*(Test|Tests|IT)$
|
||||
|
||||
# 测试实例生命周期:每个方法独立实例(默认,与 MockitoExtension 兼容)
|
||||
junit.jupiter.testinstance.lifecycle.default=per_method
|
||||
|
||||
# 显示名称:优先 @DisplayName,否则按方法名下划线分隔生成
|
||||
junit.jupiter.displayname.generator.default=org.junit.jupiter.api.DisplayNameGenerator$ReplaceUnderscores
|
||||
|
||||
# 并行执行:默认关闭(Mockito 严格 stub 不支持并行)
|
||||
junit.jupiter.execution.parallel.enabled=false
|
||||
|
||||
# 捕获标准输出供 IDE 调试
|
||||
junit.platform.output.capture.stdout=true
|
||||
junit.platform.output.capture.stderr=true
|
||||
|
||||
# ===== 按标签分组执行(在命令行或 CI 中通过 -Dgroups 指定)=====
|
||||
# 示例:
|
||||
# p0 快速核心 mvn test -Dgroups="p0"
|
||||
# p0+p1 完整 mvn test -Dgroups="p0 | p1"
|
||||
# 排除集成测试 mvn test -Dgroups="!integration"
|
||||
Reference in New Issue
Block a user