feat(phase1): 完成 Repository 测试和 Redis 会话管理

Task 2: 完成 Repository 层测试
- 实现 CaseLibraryRepositoryTest (6 个测试)
- 实现 ApiDocumentRepositoryTest (7 个测试)
- 所有 Repository 测试通过 (19/19)

Task 3: 完成 Redis 会话管理
- 创建 SessionContext 和 ToolCall 数据类
- 创建 SessionManager 接口
- 实现 RedisSessionManager (基于 Redis 的会话管理)
- 创建 SessionConfiguration (Redis 序列化配置)
- 实现 RedisSessionManagerTest (8 个测试全部通过)

测试结果:
- ApiDocumentRepositoryTest: 7/7 通过
- CaseLibraryRepositoryTest: 6/6 通过
- DiagnosisRecordRepositoryTest: 6/6 通过
- RedisSessionManagerTest: 8/8 通过
- 总计: 27/27 测试通过

功能特性:
- 会话创建、查询、更新、删除
- 会话过期时间管理
- 工具调用记录追踪
- 会话状态管理
- 基于 Redis 的分布式会话存储

Progress: 20/33 tasks completed (61%)
This commit is contained in:
zhuyongxin
2026-06-23 14:38:17 +08:00
parent 1de1e98ef8
commit 48132d297d
20 changed files with 1621 additions and 596 deletions
@@ -0,0 +1,55 @@
package org.example.config;
import com.fasterxml.jackson.annotation.JsonTypeInfo;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.fasterxml.jackson.databind.jsontype.impl.LaissezFaireSubTypeValidator;
import com.fasterxml.jackson.datatype.jsr310.JavaTimeModule;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.data.redis.connection.RedisConnectionFactory;
import org.springframework.data.redis.core.RedisTemplate;
import org.springframework.data.redis.serializer.GenericJackson2JsonRedisSerializer;
import org.springframework.data.redis.serializer.StringRedisSerializer;
/**
* Redis 会话配置
*/
@Configuration
public class SessionConfiguration {
/**
* 配置 RedisTemplate
* 使用 JSON 序列化存储会话对象
*/
@Bean
public RedisTemplate<String, Object> redisTemplate(RedisConnectionFactory connectionFactory) {
RedisTemplate<String, Object> template = new RedisTemplate<>();
template.setConnectionFactory(connectionFactory);
// 配置 ObjectMapper 支持 Java 8 时间类型
ObjectMapper objectMapper = new ObjectMapper();
objectMapper.registerModule(new JavaTimeModule());
// 启用类型信息,支持多态反序列化
objectMapper.activateDefaultTyping(
LaissezFaireSubTypeValidator.instance,
ObjectMapper.DefaultTyping.NON_FINAL,
JsonTypeInfo.As.PROPERTY
);
// 使用 JSON 序列化器
GenericJackson2JsonRedisSerializer jsonSerializer =
new GenericJackson2JsonRedisSerializer(objectMapper);
// Key 使用 String 序列化
template.setKeySerializer(new StringRedisSerializer());
template.setHashKeySerializer(new StringRedisSerializer());
// Value 使用 JSON 序列化
template.setValueSerializer(jsonSerializer);
template.setHashValueSerializer(jsonSerializer);
template.afterPropertiesSet();
return template;
}
}
@@ -0,0 +1,88 @@
package org.example.domain.model;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.NoArgsConstructor;
import java.io.Serializable;
import java.time.LocalDateTime;
import java.util.ArrayList;
import java.util.List;
/**
* 会话上下文数据类
* 存储在 Redis 中的会话数据
*/
@Data
@Builder
@NoArgsConstructor
@AllArgsConstructor
public class SessionContext implements Serializable {
private static final long serialVersionUID = 1L;
/**
* 会话ID
*/
private String sessionId;
/**
* 用户ID
*/
private String userId;
/**
* 业务ID(订单号/请求ID等)
*/
private String businessId;
/**
* 链路追踪ID
*/
private String traceId;
/**
* 会话状态(ACTIVE/COMPLETED/EXPIRED)
*/
private String status;
/**
* 工具调用历史
*/
@Builder.Default
private List<ToolCall> toolCalls = new ArrayList<>();
/**
* 会话创建时间
*/
private LocalDateTime createdAt;
/**
* 最后活跃时间
*/
private LocalDateTime lastActiveAt;
/**
* 会话过期时间(秒)
*/
private Long ttl;
/**
* 添加工具调用记录
*/
public void addToolCall(ToolCall toolCall) {
if (this.toolCalls == null) {
this.toolCalls = new ArrayList<>();
}
this.toolCalls.add(toolCall);
this.lastActiveAt = LocalDateTime.now();
}
/**
* 更新最后活跃时间
*/
public void updateLastActiveTime() {
this.lastActiveAt = LocalDateTime.now();
}
}
@@ -0,0 +1,57 @@
package org.example.domain.model;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.NoArgsConstructor;
import java.io.Serializable;
import java.time.LocalDateTime;
import java.util.Map;
/**
* 工具调用记录数据类
*/
@Data
@Builder
@NoArgsConstructor
@AllArgsConstructor
public class ToolCall implements Serializable {
private static final long serialVersionUID = 1L;
/**
* 工具名称
*/
private String toolName;
/**
* 工具参数
*/
private Map<String, Object> arguments;
/**
* 工具返回结果
*/
private String result;
/**
* 执行状态(SUCCESS/FAILED/TIMEOUT)
*/
private String status;
/**
* 错误信息(如果失败)
*/
private String errorMessage;
/**
* 执行耗时(毫秒)
*/
private Long duration;
/**
* 调用时间
*/
private LocalDateTime calledAt;
}
@@ -0,0 +1,77 @@
package org.example.service.session;
import org.example.domain.model.SessionContext;
import org.example.domain.model.ToolCall;
import java.util.Optional;
/**
* 会话管理器接口
* 负责会话的创建、读取、更新和删除
*/
public interface SessionManager {
/**
* 创建新会话
*
* @param sessionContext 会话上下文
* @param ttlSeconds 会话过期时间(秒)
* @return 会话ID
*/
String createSession(SessionContext sessionContext, long ttlSeconds);
/**
* 获取会话
*
* @param sessionId 会话ID
* @return 会话上下文(如果存在)
*/
Optional<SessionContext> getSession(String sessionId);
/**
* 更新会话
*
* @param sessionContext 会话上下文
*/
void updateSession(SessionContext sessionContext);
/**
* 删除会话
*
* @param sessionId 会话ID
*/
void deleteSession(String sessionId);
/**
* 检查会话是否存在
*
* @param sessionId 会话ID
* @return true 如果会话存在
*/
boolean exists(String sessionId);
/**
* 刷新会话过期时间
*
* @param sessionId 会话ID
* @param ttlSeconds 新的过期时间(秒)
* @return true 如果刷新成功
*/
boolean refreshSession(String sessionId, long ttlSeconds);
/**
* 添加工具调用记录到会话
*
* @param sessionId 会话ID
* @param toolCall 工具调用记录
*/
void addToolCall(String sessionId, ToolCall toolCall);
/**
* 更新会话状态
*
* @param sessionId 会话ID
* @param status 新状态
*/
void updateStatus(String sessionId, String status);
}
@@ -0,0 +1,147 @@
package org.example.service.session.impl;
import lombok.RequiredArgsConstructor;
import lombok.extern.slf4j.Slf4j;
import org.example.domain.model.SessionContext;
import org.example.domain.model.ToolCall;
import org.example.service.session.SessionManager;
import org.springframework.data.redis.core.RedisTemplate;
import org.springframework.stereotype.Service;
import java.time.LocalDateTime;
import java.util.Optional;
import java.util.concurrent.TimeUnit;
/**
* Redis 会话管理器实现
*/
@Slf4j
@Service
@RequiredArgsConstructor
public class RedisSessionManager implements SessionManager {
private static final String SESSION_KEY_PREFIX = "session:";
private final RedisTemplate<String, Object> redisTemplate;
@Override
public String createSession(SessionContext sessionContext, long ttlSeconds) {
String sessionId = sessionContext.getSessionId();
if (sessionId == null || sessionId.isEmpty()) {
throw new IllegalArgumentException("Session ID cannot be null or empty");
}
sessionContext.setCreatedAt(LocalDateTime.now());
sessionContext.setLastActiveAt(LocalDateTime.now());
sessionContext.setTtl(ttlSeconds);
sessionContext.setStatus("ACTIVE");
String key = buildKey(sessionId);
redisTemplate.opsForValue().set(key, sessionContext, ttlSeconds, TimeUnit.SECONDS);
log.info("创建会话成功: sessionId={}, ttl={}秒", sessionId, ttlSeconds);
return sessionId;
}
@Override
public Optional<SessionContext> getSession(String sessionId) {
String key = buildKey(sessionId);
Object value = redisTemplate.opsForValue().get(key);
if (value instanceof SessionContext) {
SessionContext context = (SessionContext) value;
log.debug("获取会话成功: sessionId={}", sessionId);
return Optional.of(context);
}
log.debug("会话不存在: sessionId={}", sessionId);
return Optional.empty();
}
@Override
public void updateSession(SessionContext sessionContext) {
String sessionId = sessionContext.getSessionId();
String key = buildKey(sessionId);
// 获取剩余 TTL
Long ttl = redisTemplate.getExpire(key, TimeUnit.SECONDS);
if (ttl == null || ttl <= 0) {
ttl = sessionContext.getTtl() != null ? sessionContext.getTtl() : 3600L;
}
sessionContext.setLastActiveAt(LocalDateTime.now());
redisTemplate.opsForValue().set(key, sessionContext, ttl, TimeUnit.SECONDS);
log.debug("更新会话成功: sessionId={}", sessionId);
}
@Override
public void deleteSession(String sessionId) {
String key = buildKey(sessionId);
Boolean deleted = redisTemplate.delete(key);
if (Boolean.TRUE.equals(deleted)) {
log.info("删除会话成功: sessionId={}", sessionId);
} else {
log.warn("删除会话失败,会话可能不存在: sessionId={}", sessionId);
}
}
@Override
public boolean exists(String sessionId) {
String key = buildKey(sessionId);
Boolean exists = redisTemplate.hasKey(key);
return Boolean.TRUE.equals(exists);
}
@Override
public boolean refreshSession(String sessionId, long ttlSeconds) {
String key = buildKey(sessionId);
Boolean refreshed = redisTemplate.expire(key, ttlSeconds, TimeUnit.SECONDS);
if (Boolean.TRUE.equals(refreshed)) {
log.debug("刷新会话过期时间成功: sessionId={}, newTtl={}秒", sessionId, ttlSeconds);
return true;
}
log.warn("刷新会话过期时间失败,会话可能不存在: sessionId={}", sessionId);
return false;
}
@Override
public void addToolCall(String sessionId, ToolCall toolCall) {
Optional<SessionContext> sessionOpt = getSession(sessionId);
if (sessionOpt.isPresent()) {
SessionContext context = sessionOpt.get();
context.addToolCall(toolCall);
updateSession(context);
log.debug("添加工具调用记录成功: sessionId={}, toolName={}", sessionId, toolCall.getToolName());
} else {
log.warn("会话不存在,无法添加工具调用记录: sessionId={}", sessionId);
}
}
@Override
public void updateStatus(String sessionId, String status) {
Optional<SessionContext> sessionOpt = getSession(sessionId);
if (sessionOpt.isPresent()) {
SessionContext context = sessionOpt.get();
context.setStatus(status);
updateSession(context);
log.debug("更新会话状态成功: sessionId={}, status={}", sessionId, status);
} else {
log.warn("会话不存在,无法更新状态: sessionId={}", sessionId);
}
}
/**
* 构建 Redis key
*/
private String buildKey(String sessionId) {
return SESSION_KEY_PREFIX + sessionId;
}
}
@@ -0,0 +1,64 @@
package org.example.config;
import org.junit.jupiter.api.Test;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.data.redis.core.RedisTemplate;
import javax.sql.DataSource;
import java.sql.Connection;
import static org.junit.jupiter.api.Assertions.*;
/**
* 测试 MySQL 和 Redis 连接配置
*/
@SpringBootTest
class ConnectionConfigTest {
@Autowired(required = false)
private DataSource dataSource;
@Autowired(required = false)
private RedisTemplate<String, Object> redisTemplate;
@Test
void testMySQLConnection() {
assertNotNull(dataSource, "DataSource should be configured");
try (Connection connection = dataSource.getConnection()) {
assertNotNull(connection, "Connection should not be null");
assertFalse(connection.isClosed(), "Connection should be open");
String catalog = connection.getCatalog();
System.out.println("✓ MySQL 连接成功!数据库: " + catalog);
assertEquals("superbiz_agent", catalog, "Database name should be superbiz_agent");
} catch (Exception e) {
fail("MySQL 连接失败: " + e.getMessage());
}
}
@Test
void testRedisConnection() {
assertNotNull(redisTemplate, "RedisTemplate should be configured");
try {
// 测试 PING
String testKey = "test:connection:" + System.currentTimeMillis();
String testValue = "test-value";
redisTemplate.opsForValue().set(testKey, testValue);
Object result = redisTemplate.opsForValue().get(testKey);
assertEquals(testValue, result, "Redis read/write should work");
System.out.println("✓ Redis 连接成功!读写正常");
// 清理测试数据
redisTemplate.delete(testKey);
} catch (Exception e) {
fail("Redis 连接失败: " + e.getMessage());
}
}
}
@@ -0,0 +1,71 @@
package org.example.config;
import org.junit.jupiter.api.Test;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.autoconfigure.jdbc.AutoConfigureTestDatabase;
import org.springframework.boot.test.autoconfigure.orm.jpa.DataJpaTest;
import org.springframework.test.context.TestPropertySource;
import javax.sql.DataSource;
import java.sql.Connection;
import static org.junit.jupiter.api.Assertions.*;
/**
* 单独测试 MySQL 连接(不启动完整 Spring Context)
*/
@DataJpaTest
@AutoConfigureTestDatabase(replace = AutoConfigureTestDatabase.Replace.NONE)
@TestPropertySource(properties = {
"spring.flyway.enabled=true",
"spring.jpa.hibernate.ddl-auto=validate"
})
class MySQLConnectionTest {
@Autowired
private DataSource dataSource;
@Test
void testMySQLConnection() {
assertNotNull(dataSource, "DataSource should be configured");
try (Connection connection = dataSource.getConnection()) {
assertNotNull(connection, "Connection should not be null");
assertFalse(connection.isClosed(), "Connection should be open");
String catalog = connection.getCatalog();
System.out.println("✓ MySQL 连接成功!");
System.out.println(" 数据库: " + catalog);
System.out.println(" URL: " + connection.getMetaData().getURL());
assertEquals("superbiz_agent", catalog, "Database name should be superbiz_agent");
} catch (Exception e) {
fail("MySQL 连接失败: " + e.getMessage());
}
}
@Test
void testFlywayMigration() {
// 如果能到这里,说明 Flyway 迁移成功
System.out.println("✓ Flyway 迁移已执行");
try (Connection connection = dataSource.getConnection()) {
var metadata = connection.getMetaData();
var tables = metadata.getTables(null, null, "%", new String[]{"TABLE"});
int tableCount = 0;
System.out.println(" 已创建的表:");
while (tables.next()) {
String tableName = tables.getString("TABLE_NAME");
System.out.println(" - " + tableName);
tableCount++;
}
assertTrue(tableCount >= 3, "至少应该有 3 张表 (diagnosis_record, case_library, api_document)");
} catch (Exception e) {
fail("验证表结构失败: " + e.getMessage());
}
}
}
@@ -0,0 +1,55 @@
package org.example.config;
import org.junit.jupiter.api.Test;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.data.redis.core.RedisTemplate;
import org.springframework.test.context.TestPropertySource;
import static org.junit.jupiter.api.Assertions.*;
/**
* 单独测试 Redis 连接
* 禁用 Milvus 以避免启动失败
*/
@SpringBootTest
@TestPropertySource(properties = {
"spring.autoconfigure.exclude=org.example.config.MilvusConfig"
})
class RedisConnectionTest {
@Autowired(required = false)
private RedisTemplate<String, Object> redisTemplate;
@Test
void testRedisConnection() {
assertNotNull(redisTemplate, "RedisTemplate should be configured");
try {
String testKey = "test:connection:" + System.currentTimeMillis();
String testValue = "test-value";
// 测试写入
redisTemplate.opsForValue().set(testKey, testValue);
System.out.println("✓ Redis 写入成功");
// 测试读取
Object result = redisTemplate.opsForValue().get(testKey);
assertEquals(testValue, result, "Redis read should return the written value");
System.out.println("✓ Redis 读取成功");
// 测试删除
redisTemplate.delete(testKey);
Object deleted = redisTemplate.opsForValue().get(testKey);
assertNull(deleted, "Deleted key should return null");
System.out.println("✓ Redis 删除成功");
System.out.println("\n✓ Redis 连接测试通过!");
System.out.println(" 服务器: 119.29.78.52:6379");
System.out.println(" 读写删操作正常");
} catch (Exception e) {
fail("Redis 连接失败: " + e.getMessage());
}
}
}
@@ -0,0 +1,50 @@
package org.example.config;
import org.junit.jupiter.api.Test;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.data.redis.core.StringRedisTemplate;
import static org.junit.jupiter.api.Assertions.*;
/**
* Redis 连接测试
* 注意:如果 Milvus 未启动,此测试会因 Spring Context 加载失败而失败
* 使用 @DataRedisTest 可以避免加载完整 Context,但需要额外配置
*/
@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.NONE)
class SimpleRedisTest {
@Autowired(required = false)
private StringRedisTemplate stringRedisTemplate;
@Test
void testRedisBasicOperations() {
// 如果能到这里,说明 Spring Context 加载成功,Redis 配置正确
assertNotNull(stringRedisTemplate, "StringRedisTemplate should be configured");
String testKey = "test:simple:" + System.currentTimeMillis();
String testValue = "hello-redis";
try {
// 写入
stringRedisTemplate.opsForValue().set(testKey, testValue);
System.out.println("✓ Redis 写入成功: " + testKey + " = " + testValue);
// 读取
String result = stringRedisTemplate.opsForValue().get(testKey);
assertEquals(testValue, result);
System.out.println("✓ Redis 读取成功: " + result);
// 删除
Boolean deleted = stringRedisTemplate.delete(testKey);
assertTrue(deleted != null && deleted);
System.out.println("✓ Redis 删除成功");
System.out.println("\n✅ Redis 连接正常,读写删操作成功");
} catch (Exception e) {
fail("Redis 操作失败: " + e.getMessage());
}
}
}
@@ -0,0 +1,179 @@
package org.example.repository;
import com.superbiz.agent.domain.enums.FaultCategory;
import org.example.domain.entity.ApiDocument;
import org.junit.jupiter.api.Test;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.autoconfigure.jdbc.AutoConfigureTestDatabase;
import org.springframework.boot.test.autoconfigure.orm.jpa.DataJpaTest;
import org.springframework.data.domain.Page;
import org.springframework.data.domain.PageRequest;
import org.springframework.test.context.TestPropertySource;
import java.time.LocalDateTime;
import java.util.List;
import java.util.Optional;
import java.util.UUID;
import static org.junit.jupiter.api.Assertions.*;
/**
* ApiDocumentRepository 单元测试
*/
@DataJpaTest
@AutoConfigureTestDatabase(replace = AutoConfigureTestDatabase.Replace.NONE)
@TestPropertySource(properties = {
"spring.flyway.enabled=true",
"spring.jpa.hibernate.ddl-auto=validate",
"spring.jpa.show-sql=false"
})
class ApiDocumentRepositoryTest {
@Autowired
private ApiDocumentRepository repository;
@Test
void testSaveAndFindById() {
ApiDocument doc = ApiDocument.builder()
.docId(UUID.randomUUID().toString())
.fileName("api-spec.md")
.filePath("/uploads/api-spec.md")
.fileHash("abc123hash")
.fileSize(1024L)
.faultCategory(FaultCategory.EXTERNAL_API)
.faultSource("广东")
.apiName("查询接口")
.status("PENDING")
.chunkCount(0)
.build();
ApiDocument saved = repository.save(doc);
assertNotNull(saved.getId());
assertNotNull(saved.getCreatedAt());
System.out.println("✓ 保存文档成功,ID: " + saved.getId());
Optional<ApiDocument> found = repository.findById(saved.getId());
assertTrue(found.isPresent());
assertEquals("api-spec.md", found.get().getFileName());
System.out.println("✓ 查询文档成功");
}
@Test
void testFindByDocId() {
String docId = UUID.randomUUID().toString();
ApiDocument doc = ApiDocument.builder()
.docId(docId)
.fileName("test-doc.txt")
.status("INDEXED")
.build();
repository.save(doc);
Optional<ApiDocument> found = repository.findByDocId(docId);
assertTrue(found.isPresent());
assertEquals(docId, found.get().getDocId());
System.out.println("✓ 根据 docId 查询成功");
}
@Test
void testFindByFileHash() {
String fileHash = "unique-hash-" + System.currentTimeMillis();
ApiDocument doc = ApiDocument.builder()
.docId(UUID.randomUUID().toString())
.fileName("duplicate-check.md")
.fileHash(fileHash)
.build();
repository.save(doc);
Optional<ApiDocument> found = repository.findByFileHash(fileHash);
assertTrue(found.isPresent());
assertEquals(fileHash, found.get().getFileHash());
System.out.println("✓ 根据 fileHash 查询成功(去重检测)");
}
@Test
void testFindByStatus() {
ApiDocument doc1 = ApiDocument.builder()
.docId(UUID.randomUUID().toString())
.fileName("pending-1.md")
.status("PENDING")
.build();
ApiDocument doc2 = ApiDocument.builder()
.docId(UUID.randomUUID().toString())
.fileName("pending-2.md")
.status("PENDING")
.build();
repository.save(doc1);
repository.save(doc2);
List<ApiDocument> results = repository.findByStatus("PENDING");
assertFalse(results.isEmpty());
assertTrue(results.size() >= 2);
System.out.println("✓ 根据状态查询成功,找到 " + results.size() + " 条 PENDING 文档");
}
@Test
void testFindByStatusWithPagination() {
// 创建多个文档
for (int i = 0; i < 5; i++) {
ApiDocument doc = ApiDocument.builder()
.docId(UUID.randomUUID().toString())
.fileName("doc-" + i + ".txt")
.status("INDEXED")
.chunkCount(10)
.build();
repository.save(doc);
}
Page<ApiDocument> page = repository.findByStatus("INDEXED", PageRequest.of(0, 3));
assertFalse(page.isEmpty());
assertTrue(page.getContent().size() <= 3);
System.out.println("✓ 分页查询成功,本页 " + page.getContent().size() + " 条,总计 " + page.getTotalElements() + " 条");
}
@Test
void testUpdateDocumentStatus() {
ApiDocument doc = ApiDocument.builder()
.docId(UUID.randomUUID().toString())
.fileName("indexing-doc.md")
.status("PENDING")
.chunkCount(0)
.build();
ApiDocument saved = repository.save(doc);
Long id = saved.getId();
// 更新状态
saved.setStatus("INDEXED");
saved.setChunkCount(15);
saved.setIndexedAt(LocalDateTime.now());
repository.save(saved);
Optional<ApiDocument> updated = repository.findById(id);
assertTrue(updated.isPresent());
assertEquals("INDEXED", updated.get().getStatus());
assertEquals(15, updated.get().getChunkCount());
assertNotNull(updated.get().getIndexedAt());
System.out.println("✓ 更新文档状态成功");
}
@Test
void testFindByFaultSource() {
ApiDocument doc = ApiDocument.builder()
.docId(UUID.randomUUID().toString())
.fileName("guangdong-api.md")
.faultSource("广东")
.faultCategory(FaultCategory.EXTERNAL_API)
.build();
repository.save(doc);
List<ApiDocument> results = repository.findByFaultSource("广东");
assertFalse(results.isEmpty());
System.out.println("✓ 根据故障源查询成功,找到 " + results.size() + " 条文档");
}
}
@@ -0,0 +1,173 @@
package org.example.repository;
import com.superbiz.agent.domain.enums.FaultCategory;
import com.superbiz.agent.domain.enums.SourceType;
import org.example.domain.entity.CaseLibrary;
import org.junit.jupiter.api.Test;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.autoconfigure.jdbc.AutoConfigureTestDatabase;
import org.springframework.boot.test.autoconfigure.orm.jpa.DataJpaTest;
import org.springframework.data.domain.Page;
import org.springframework.data.domain.PageRequest;
import org.springframework.test.context.TestPropertySource;
import java.util.List;
import java.util.Optional;
import java.util.UUID;
import static org.junit.jupiter.api.Assertions.*;
/**
* CaseLibraryRepository 单元测试
*/
@DataJpaTest
@AutoConfigureTestDatabase(replace = AutoConfigureTestDatabase.Replace.NONE)
@TestPropertySource(properties = {
"spring.flyway.enabled=true",
"spring.jpa.hibernate.ddl-auto=validate",
"spring.jpa.show-sql=false"
})
class CaseLibraryRepositoryTest {
@Autowired
private CaseLibraryRepository repository;
@Test
void testSaveAndFindById() {
CaseLibrary caseLib = CaseLibrary.builder()
.caseId(UUID.randomUUID().toString())
.title("接口超时案例")
.rootCause("网络延迟导致接口超时")
.solution("增加超时时间和重试机制")
.faultCategory(FaultCategory.EXTERNAL_API)
.errorCode("40003")
.sourceType(SourceType.AUTO)
.build();
CaseLibrary saved = repository.save(caseLib);
assertNotNull(saved.getId());
assertNotNull(saved.getCreatedAt());
System.out.println("✓ 保存案例成功,ID: " + saved.getId());
Optional<CaseLibrary> found = repository.findById(saved.getId());
assertTrue(found.isPresent());
assertEquals("接口超时案例", found.get().getTitle());
System.out.println("✓ 查询案例成功");
}
@Test
void testFindByCaseId() {
String caseId = UUID.randomUUID().toString();
CaseLibrary caseLib = CaseLibrary.builder()
.caseId(caseId)
.title("数据库死锁案例")
.rootCause("并发更新导致死锁")
.solution("优化事务粒度")
.faultCategory(FaultCategory.DATABASE)
.build();
repository.save(caseLib);
Optional<CaseLibrary> found = repository.findByCaseId(caseId);
assertTrue(found.isPresent());
assertEquals(caseId, found.get().getCaseId());
System.out.println("✓ 根据 caseId 查询成功");
}
@Test
void testFindByFaultCategoryAndErrorCode() {
CaseLibrary case1 = CaseLibrary.builder()
.caseId(UUID.randomUUID().toString())
.title("案例1")
.rootCause("原因1")
.solution("方案1")
.faultCategory(FaultCategory.EXTERNAL_API)
.errorCode("40003")
.build();
CaseLibrary case2 = CaseLibrary.builder()
.caseId(UUID.randomUUID().toString())
.title("案例2")
.rootCause("原因2")
.solution("方案2")
.faultCategory(FaultCategory.EXTERNAL_API)
.errorCode("40003")
.build();
repository.save(case1);
repository.save(case2);
List<CaseLibrary> results = repository.findByFaultCategoryAndErrorCode(
FaultCategory.EXTERNAL_API, "40003");
assertFalse(results.isEmpty());
assertTrue(results.size() >= 2);
System.out.println("✓ 根据故障类别和错误码查询成功,找到 " + results.size() + " 条案例");
}
@Test
void testFindBySourceType() {
CaseLibrary autoCase = CaseLibrary.builder()
.caseId(UUID.randomUUID().toString())
.title("自动生成案例")
.rootCause("自动分析")
.solution("自动方案")
.sourceType(SourceType.AUTO)
.build();
repository.save(autoCase);
Page<CaseLibrary> results = repository.findBySourceType(
SourceType.AUTO, PageRequest.of(0, 10));
assertFalse(results.isEmpty());
System.out.println("✓ 根据来源类型查询成功,找到 " + results.getTotalElements() + " 条案例");
}
@Test
void testUpdateReferenceCount() {
CaseLibrary caseLib = CaseLibrary.builder()
.caseId(UUID.randomUUID().toString())
.title("热门案例")
.rootCause("常见问题")
.solution("标准方案")
.referenceCount(0)
.build();
CaseLibrary saved = repository.save(caseLib);
Long id = saved.getId();
// 模拟被引用
saved.setReferenceCount(saved.getReferenceCount() + 1);
repository.save(saved);
Optional<CaseLibrary> updated = repository.findById(id);
assertTrue(updated.isPresent());
assertEquals(1, updated.get().getReferenceCount());
System.out.println("✓ 更新引用次数成功");
}
@Test
void testFindTopByReferenceCount() {
// 创建一些案例
for (int i = 0; i < 3; i++) {
CaseLibrary caseLib = CaseLibrary.builder()
.caseId(UUID.randomUUID().toString())
.title("案例 " + i)
.rootCause("原因")
.solution("方案")
.referenceCount(i * 10)
.build();
repository.save(caseLib);
}
List<CaseLibrary> topCases = repository.findTop10ByOrderByReferenceCountDesc();
assertFalse(topCases.isEmpty());
// 验证排序(第一个引用次数应该最高)
if (topCases.size() >= 2) {
assertTrue(topCases.get(0).getReferenceCount() >= topCases.get(1).getReferenceCount());
}
System.out.println("✓ 查询热门案例成功,返回 " + topCases.size() + " 条");
}
}
@@ -0,0 +1,244 @@
package org.example.service.session;
import org.example.domain.model.SessionContext;
import org.example.domain.model.ToolCall;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.test.context.TestPropertySource;
import java.time.LocalDateTime;
import java.util.HashMap;
import java.util.Map;
import java.util.Optional;
import java.util.UUID;
import static org.junit.jupiter.api.Assertions.*;
/**
* RedisSessionManager 单元测试
*/
@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.NONE)
@TestPropertySource(properties = {
"spring.redis.host=119.29.78.52",
"spring.redis.port=6379",
"spring.redis.password=!Fucker123.."
})
class RedisSessionManagerTest {
@Autowired
private SessionManager sessionManager;
private String testSessionId;
@BeforeEach
void setUp() {
testSessionId = "test-session-" + UUID.randomUUID().toString();
}
@Test
void testCreateAndGetSession() {
// 创建会话
SessionContext context = SessionContext.builder()
.sessionId(testSessionId)
.userId("user-123")
.businessId("order-456")
.traceId("trace-789")
.build();
String sessionId = sessionManager.createSession(context, 300); // 5分钟
assertNotNull(sessionId);
assertEquals(testSessionId, sessionId);
System.out.println("✓ 创建会话成功: " + sessionId);
// 获取会话
Optional<SessionContext> retrieved = sessionManager.getSession(testSessionId);
assertTrue(retrieved.isPresent());
assertEquals("user-123", retrieved.get().getUserId());
assertEquals("ACTIVE", retrieved.get().getStatus());
assertNotNull(retrieved.get().getCreatedAt());
System.out.println("✓ 获取会话成功");
// 清理
sessionManager.deleteSession(testSessionId);
}
@Test
void testUpdateSession() {
// 创建会话
SessionContext context = SessionContext.builder()
.sessionId(testSessionId)
.userId("user-update")
.status("ACTIVE")
.build();
sessionManager.createSession(context, 300);
// 获取并更新
Optional<SessionContext> retrieved = sessionManager.getSession(testSessionId);
assertTrue(retrieved.isPresent());
SessionContext toUpdate = retrieved.get();
toUpdate.setStatus("COMPLETED");
toUpdate.setBusinessId("updated-business-id");
sessionManager.updateSession(toUpdate);
// 验证更新
Optional<SessionContext> updated = sessionManager.getSession(testSessionId);
assertTrue(updated.isPresent());
assertEquals("COMPLETED", updated.get().getStatus());
assertEquals("updated-business-id", updated.get().getBusinessId());
System.out.println("✓ 更新会话成功");
// 清理
sessionManager.deleteSession(testSessionId);
}
@Test
void testDeleteSession() {
// 创建会话
SessionContext context = SessionContext.builder()
.sessionId(testSessionId)
.userId("user-delete")
.build();
sessionManager.createSession(context, 300);
assertTrue(sessionManager.exists(testSessionId));
// 删除会话
sessionManager.deleteSession(testSessionId);
assertFalse(sessionManager.exists(testSessionId));
System.out.println("✓ 删除会话成功");
}
@Test
void testExists() {
assertFalse(sessionManager.exists(testSessionId));
SessionContext context = SessionContext.builder()
.sessionId(testSessionId)
.userId("user-exists")
.build();
sessionManager.createSession(context, 300);
assertTrue(sessionManager.exists(testSessionId));
System.out.println("✓ 会话存在性检查成功");
// 清理
sessionManager.deleteSession(testSessionId);
}
@Test
void testRefreshSession() {
// 创建会话
SessionContext context = SessionContext.builder()
.sessionId(testSessionId)
.userId("user-refresh")
.build();
sessionManager.createSession(context, 60); // 1分钟
// 刷新过期时间
boolean refreshed = sessionManager.refreshSession(testSessionId, 600); // 延长到10分钟
assertTrue(refreshed);
assertTrue(sessionManager.exists(testSessionId));
System.out.println("✓ 刷新会话过期时间成功");
// 清理
sessionManager.deleteSession(testSessionId);
}
@Test
void testAddToolCall() {
// 创建会话
SessionContext context = SessionContext.builder()
.sessionId(testSessionId)
.userId("user-toolcall")
.build();
sessionManager.createSession(context, 300);
// 添加工具调用记录
Map<String, Object> args = new HashMap<>();
args.put("query", "test query");
args.put("limit", 10);
ToolCall toolCall = ToolCall.builder()
.toolName("search_documents")
.arguments(args)
.result("found 5 documents")
.status("SUCCESS")
.duration(150L)
.calledAt(LocalDateTime.now())
.build();
sessionManager.addToolCall(testSessionId, toolCall);
// 验证工具调用已添加
Optional<SessionContext> retrieved = sessionManager.getSession(testSessionId);
assertTrue(retrieved.isPresent());
assertFalse(retrieved.get().getToolCalls().isEmpty());
assertEquals(1, retrieved.get().getToolCalls().size());
assertEquals("search_documents", retrieved.get().getToolCalls().get(0).getToolName());
System.out.println("✓ 添加工具调用记录成功");
// 清理
sessionManager.deleteSession(testSessionId);
}
@Test
void testUpdateStatus() {
// 创建会话
SessionContext context = SessionContext.builder()
.sessionId(testSessionId)
.userId("user-status")
.status("ACTIVE")
.build();
sessionManager.createSession(context, 300);
// 更新状态
sessionManager.updateStatus(testSessionId, "COMPLETED");
// 验证状态已更新
Optional<SessionContext> retrieved = sessionManager.getSession(testSessionId);
assertTrue(retrieved.isPresent());
assertEquals("COMPLETED", retrieved.get().getStatus());
System.out.println("✓ 更新会话状态成功");
// 清理
sessionManager.deleteSession(testSessionId);
}
@Test
void testMultipleToolCalls() {
// 创建会话
SessionContext context = SessionContext.builder()
.sessionId(testSessionId)
.userId("user-multi-tools")
.build();
sessionManager.createSession(context, 300);
// 添加多个工具调用
for (int i = 0; i < 3; i++) {
ToolCall toolCall = ToolCall.builder()
.toolName("tool_" + i)
.status("SUCCESS")
.calledAt(LocalDateTime.now())
.build();
sessionManager.addToolCall(testSessionId, toolCall);
}
// 验证所有工具调用
Optional<SessionContext> retrieved = sessionManager.getSession(testSessionId);
assertTrue(retrieved.isPresent());
assertEquals(3, retrieved.get().getToolCalls().size());
System.out.println("✓ 添加多个工具调用记录成功");
// 清理
sessionManager.deleteSession(testSessionId);
}
}