feat(session): 会话存储体系实现 & Chat多Agent路由
- 新增诊断会话(diagnosis_session/agent_step/tool_invocation)三表 - AgentLoggingHook 持久化 agent_step,记录决策链和耗时 - LookupKnowledgeTool 写入 tool_invocation,记录L0/L1检索质量 - TokenTrackingChatModel 捕获真实token用量 - Chat接口支持意图路由:简单问题单Agent,复杂问题多Agent(Planner+Executor) - Prompt外置到 src/main/resources/prompts/ - 删除旧 diagnosis_record 表及相关文件 - 新增SessionContextHolder(ThreadLocal传递sessionId) - QuestionComplexity 复杂度判断工具 - 测试覆盖三张新表的Repository
This commit is contained in:
@@ -0,0 +1,61 @@
|
||||
package com.superbiz.agent.repository;
|
||||
|
||||
import com.superbiz.agent.domain.entity.AgentStep;
|
||||
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 java.util.List;
|
||||
import java.util.UUID;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
|
||||
/**
|
||||
* AgentStepRepository 单元测试
|
||||
*/
|
||||
@DataJpaTest
|
||||
@AutoConfigureTestDatabase(replace = AutoConfigureTestDatabase.Replace.NONE)
|
||||
@TestPropertySource(properties = {
|
||||
"spring.flyway.enabled=true",
|
||||
"spring.jpa.hibernate.ddl-auto=validate",
|
||||
"spring.jpa.show-sql=true"
|
||||
})
|
||||
class AgentStepRepositoryTest {
|
||||
|
||||
@Autowired
|
||||
private AgentStepRepository repository;
|
||||
|
||||
@Test
|
||||
void testSaveAndFindBySessionId() {
|
||||
String sessionId = UUID.randomUUID().toString().substring(0, 8);
|
||||
|
||||
AgentStep step0 = AgentStep.builder()
|
||||
.sessionId(sessionId)
|
||||
.stepIndex(0)
|
||||
.agentName("planner")
|
||||
.hasToolCall(true)
|
||||
.durationMs(500)
|
||||
.build();
|
||||
repository.save(step0);
|
||||
|
||||
AgentStep step1 = AgentStep.builder()
|
||||
.sessionId(sessionId)
|
||||
.stepIndex(1)
|
||||
.agentName("executor")
|
||||
.hasToolCall(false)
|
||||
.durationMs(300)
|
||||
.build();
|
||||
repository.save(step1);
|
||||
|
||||
List<AgentStep> steps = repository.findBySessionIdOrderByStepIndex(sessionId);
|
||||
assertEquals(2, steps.size());
|
||||
assertEquals("planner", steps.get(0).getAgentName());
|
||||
assertEquals("executor", steps.get(1).getAgentName());
|
||||
assertEquals(500, steps.get(0).getDurationMs());
|
||||
|
||||
int count = repository.countBySessionId(sessionId);
|
||||
assertEquals(2, count);
|
||||
}
|
||||
}
|
||||
@@ -1,173 +0,0 @@
|
||||
package com.superbiz.agent.repository;
|
||||
|
||||
import com.superbiz.agent.domain.enums.DiagnosisStatus;
|
||||
import com.superbiz.agent.domain.enums.FaultCategory;
|
||||
import com.superbiz.agent.domain.entity.DiagnosisRecord;
|
||||
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 java.util.List;
|
||||
import java.util.Optional;
|
||||
import java.util.UUID;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
|
||||
/**
|
||||
* DiagnosisRecordRepository 单元测试
|
||||
*/
|
||||
@DataJpaTest
|
||||
@AutoConfigureTestDatabase(replace = AutoConfigureTestDatabase.Replace.NONE)
|
||||
@TestPropertySource(properties = {
|
||||
"spring.flyway.enabled=true",
|
||||
"spring.jpa.hibernate.ddl-auto=validate",
|
||||
"spring.jpa.show-sql=true"
|
||||
})
|
||||
class DiagnosisRecordRepositoryTest {
|
||||
|
||||
@Autowired
|
||||
private DiagnosisRecordRepository repository;
|
||||
|
||||
@Test
|
||||
void testSaveAndFindById() {
|
||||
// 创建测试数据
|
||||
DiagnosisRecord record = DiagnosisRecord.builder()
|
||||
.diagnosisId(UUID.randomUUID().toString())
|
||||
.sessionId("session-001")
|
||||
.businessId("order-12345")
|
||||
.traceId("trace-abc123")
|
||||
.faultCategory(FaultCategory.API)
|
||||
.faultSource("广东")
|
||||
.faultTarget("http://api.example.com/query")
|
||||
.errorCode("40003")
|
||||
.errorMessage("接口超时")
|
||||
.status(DiagnosisStatus.SUCCESS)
|
||||
.confidence(85)
|
||||
.duration(1500)
|
||||
.build();
|
||||
|
||||
// 保存
|
||||
DiagnosisRecord saved = repository.save(record);
|
||||
assertNotNull(saved.getId());
|
||||
assertNotNull(saved.getCreatedAt());
|
||||
System.out.println("✓ 保存成功,ID: " + saved.getId());
|
||||
|
||||
// 查询
|
||||
Optional<DiagnosisRecord> found = repository.findById(saved.getId());
|
||||
assertTrue(found.isPresent());
|
||||
assertEquals("order-12345", found.get().getBusinessId());
|
||||
System.out.println("✓ 根据 ID 查询成功");
|
||||
}
|
||||
|
||||
@Test
|
||||
void testFindByDiagnosisId() {
|
||||
String diagnosisId = UUID.randomUUID().toString();
|
||||
DiagnosisRecord record = DiagnosisRecord.builder()
|
||||
.diagnosisId(diagnosisId)
|
||||
.businessId("order-test-001")
|
||||
.faultCategory(FaultCategory.API)
|
||||
.status(DiagnosisStatus.PENDING)
|
||||
.build();
|
||||
|
||||
repository.save(record);
|
||||
|
||||
Optional<DiagnosisRecord> found = repository.findByDiagnosisId(diagnosisId);
|
||||
assertTrue(found.isPresent());
|
||||
assertEquals(diagnosisId, found.get().getDiagnosisId());
|
||||
System.out.println("✓ 根据 diagnosisId 查询成功");
|
||||
}
|
||||
|
||||
@Test
|
||||
void testFindByFaultCategoryAndErrorCode() {
|
||||
// 创建测试数据
|
||||
DiagnosisRecord record1 = DiagnosisRecord.builder()
|
||||
.diagnosisId(UUID.randomUUID().toString())
|
||||
.faultCategory(FaultCategory.API)
|
||||
.errorCode("40003")
|
||||
.status(DiagnosisStatus.SUCCESS)
|
||||
.build();
|
||||
|
||||
DiagnosisRecord record2 = DiagnosisRecord.builder()
|
||||
.diagnosisId(UUID.randomUUID().toString())
|
||||
.faultCategory(FaultCategory.API)
|
||||
.errorCode("40003")
|
||||
.status(DiagnosisStatus.FAILED)
|
||||
.build();
|
||||
|
||||
repository.save(record1);
|
||||
repository.save(record2);
|
||||
|
||||
// 查询
|
||||
List<DiagnosisRecord> results = repository.findByFaultCategoryAndErrorCode(
|
||||
FaultCategory.API, "40003");
|
||||
|
||||
assertFalse(results.isEmpty());
|
||||
assertTrue(results.size() >= 2);
|
||||
System.out.println("✓ 根据故障类别和错误码查询成功,找到 " + results.size() + " 条记录");
|
||||
}
|
||||
|
||||
@Test
|
||||
void testFindByStatus() {
|
||||
DiagnosisRecord record = DiagnosisRecord.builder()
|
||||
.diagnosisId(UUID.randomUUID().toString())
|
||||
.status(DiagnosisStatus.RUNNING)
|
||||
.faultCategory(FaultCategory.API)
|
||||
.build();
|
||||
|
||||
repository.save(record);
|
||||
|
||||
List<DiagnosisRecord> results = repository.findByStatus(DiagnosisStatus.RUNNING);
|
||||
assertFalse(results.isEmpty());
|
||||
System.out.println("✓ 根据状态查询成功,找到 " + results.size() + " 条 RUNNING 记录");
|
||||
}
|
||||
|
||||
@Test
|
||||
void testUpdateRecord() {
|
||||
// 创建并保存
|
||||
DiagnosisRecord record = DiagnosisRecord.builder()
|
||||
.diagnosisId(UUID.randomUUID().toString())
|
||||
.status(DiagnosisStatus.PENDING)
|
||||
.confidence(0)
|
||||
.build();
|
||||
|
||||
DiagnosisRecord saved = repository.save(record);
|
||||
Long id = saved.getId();
|
||||
|
||||
// 更新
|
||||
saved.setStatus(DiagnosisStatus.SUCCESS);
|
||||
saved.setConfidence(90);
|
||||
saved.setRootCause("接口超时导致");
|
||||
saved.setSolution("增加重试机制");
|
||||
|
||||
repository.save(saved);
|
||||
|
||||
// 验证更新
|
||||
Optional<DiagnosisRecord> updated = repository.findById(id);
|
||||
assertTrue(updated.isPresent());
|
||||
assertEquals(DiagnosisStatus.SUCCESS, updated.get().getStatus());
|
||||
assertEquals(90, updated.get().getConfidence());
|
||||
assertNotNull(updated.get().getUpdatedAt());
|
||||
System.out.println("✓ 更新记录成功");
|
||||
}
|
||||
|
||||
@Test
|
||||
void testDeleteRecord() {
|
||||
DiagnosisRecord record = DiagnosisRecord.builder()
|
||||
.diagnosisId(UUID.randomUUID().toString())
|
||||
.status(DiagnosisStatus.PENDING)
|
||||
.build();
|
||||
|
||||
DiagnosisRecord saved = repository.save(record);
|
||||
Long id = saved.getId();
|
||||
|
||||
// 删除
|
||||
repository.deleteById(id);
|
||||
|
||||
// 验证删除
|
||||
Optional<DiagnosisRecord> deleted = repository.findById(id);
|
||||
assertFalse(deleted.isPresent());
|
||||
System.out.println("✓ 删除记录成功");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package com.superbiz.agent.repository;
|
||||
|
||||
import com.superbiz.agent.domain.entity.DiagnosisSession;
|
||||
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 java.util.Optional;
|
||||
import java.util.UUID;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
|
||||
/**
|
||||
* DiagnosisSessionRepository 单元测试
|
||||
*/
|
||||
@DataJpaTest
|
||||
@AutoConfigureTestDatabase(replace = AutoConfigureTestDatabase.Replace.NONE)
|
||||
@TestPropertySource(properties = {
|
||||
"spring.flyway.enabled=true",
|
||||
"spring.jpa.hibernate.ddl-auto=validate",
|
||||
"spring.jpa.show-sql=true"
|
||||
})
|
||||
class DiagnosisSessionRepositoryTest {
|
||||
|
||||
@Autowired
|
||||
private DiagnosisSessionRepository repository;
|
||||
|
||||
@Test
|
||||
void testSaveAndFindBySessionId() {
|
||||
String sessionId = UUID.randomUUID().toString().substring(0, 8);
|
||||
DiagnosisSession session = DiagnosisSession.builder()
|
||||
.sessionId(sessionId)
|
||||
.query("测试查询")
|
||||
.status("RUNNING")
|
||||
.agentFlow("CHAT")
|
||||
.build();
|
||||
|
||||
DiagnosisSession saved = repository.save(session);
|
||||
assertNotNull(saved.getId());
|
||||
assertEquals(sessionId, saved.getSessionId());
|
||||
|
||||
Optional<DiagnosisSession> found = repository.findBySessionId(sessionId);
|
||||
assertTrue(found.isPresent());
|
||||
assertEquals("测试查询", found.get().getQuery());
|
||||
assertEquals("CHAT", found.get().getAgentFlow());
|
||||
}
|
||||
|
||||
@Test
|
||||
void testUpdateStatus() {
|
||||
String sessionId = UUID.randomUUID().toString().substring(0, 8);
|
||||
DiagnosisSession session = DiagnosisSession.builder()
|
||||
.sessionId(sessionId)
|
||||
.query("更新测试")
|
||||
.status("RUNNING")
|
||||
.agentFlow("AI_OPS")
|
||||
.build();
|
||||
|
||||
DiagnosisSession saved = repository.save(session);
|
||||
saved.setStatus("SUCCESS");
|
||||
saved.setTotalDurationMs(1500);
|
||||
repository.save(saved);
|
||||
|
||||
DiagnosisSession updated = repository.findBySessionId(sessionId).orElseThrow();
|
||||
assertEquals("SUCCESS", updated.getStatus());
|
||||
assertEquals(1500, updated.getTotalDurationMs());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,64 @@
|
||||
package com.superbiz.agent.repository;
|
||||
|
||||
import com.superbiz.agent.domain.entity.ToolInvocation;
|
||||
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 java.util.List;
|
||||
import java.util.UUID;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
|
||||
/**
|
||||
* ToolInvocationRepository 单元测试
|
||||
*/
|
||||
@DataJpaTest
|
||||
@AutoConfigureTestDatabase(replace = AutoConfigureTestDatabase.Replace.NONE)
|
||||
@TestPropertySource(properties = {
|
||||
"spring.flyway.enabled=true",
|
||||
"spring.jpa.hibernate.ddl-auto=validate",
|
||||
"spring.jpa.show-sql=true"
|
||||
})
|
||||
class ToolInvocationRepositoryTest {
|
||||
|
||||
@Autowired
|
||||
private ToolInvocationRepository repository;
|
||||
|
||||
@Test
|
||||
void testSaveAndFindBySessionId() {
|
||||
String sessionId = UUID.randomUUID().toString().substring(0, 8);
|
||||
|
||||
ToolInvocation inv1 = ToolInvocation.builder()
|
||||
.sessionId(sessionId)
|
||||
.toolName("lookup_knowledge")
|
||||
.inputParams("{\"query\":\"ERR_TIMEOUT\"}")
|
||||
.retrievalLayer("L0")
|
||||
.l0MatchCount(1)
|
||||
.durationMs(50)
|
||||
.success(true)
|
||||
.build();
|
||||
repository.save(inv1);
|
||||
|
||||
ToolInvocation inv2 = ToolInvocation.builder()
|
||||
.sessionId(sessionId)
|
||||
.toolName("queryPrometheusAlerts")
|
||||
.inputParams("{\"metric\":\"cpu_usage\"}")
|
||||
.durationMs(200)
|
||||
.success(true)
|
||||
.build();
|
||||
repository.save(inv2);
|
||||
|
||||
List<ToolInvocation> bySession = repository.findBySessionId(sessionId);
|
||||
assertEquals(2, bySession.size());
|
||||
|
||||
List<ToolInvocation> byTool = repository.findByToolName("lookup_knowledge");
|
||||
assertFalse(byTool.isEmpty());
|
||||
|
||||
List<ToolInvocation> byBoth = repository.findBySessionIdAndToolName(sessionId, "lookup_knowledge");
|
||||
assertEquals(1, byBoth.size());
|
||||
assertEquals("L0", byBoth.get(0).getRetrievalLayer());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user