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:
zhuyongxin
2026-06-26 16:22:05 +08:00
parent a74ccea5be
commit 0d9cce75f9
33 changed files with 1941 additions and 470 deletions
@@ -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());
}
}