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:
nojava
2026-05-24 03:36:31 +08:00
parent 9cf811d8c1
commit 6c8b31e38a
58 changed files with 2693 additions and 556 deletions
@@ -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();
}
}
@@ -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
}
@@ -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);
}
@@ -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);
}
@@ -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();
}
}
@@ -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();
}
@@ -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());
}
}
@@ -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());
}
}
@@ -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();
@@ -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));
}
@@ -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);
@@ -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 用于断言传入参数 / 抛出异常。
*/
@@ -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("JavaSpring Boot,微服务"));
.thenReturn(TestDataFactory.modelResponse("JavaSpring 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
)
);
}
}
@@ -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"