feat(trace): isolate chat runs

This commit is contained in:
zhuyongxin
2026-07-10 19:02:04 +08:00
parent 6fdbd34bab
commit 26d5529280
21 changed files with 610 additions and 134 deletions
@@ -1,6 +1,6 @@
# ISS-010 同 session 多轮诊断 Trace 隔离
**状态**:方案已确认,待 OpenSpec
**状态**:OpenSpec 已创建,Phase 1 已完成,Phase 2 进行中
**严重程度**:高
**发现时间**:2026-07-10
**来源**:同一 `sessionId` 多轮 Chat E2E 验证
@@ -439,8 +439,6 @@ tool_invocation(run_id, id)
- 增加 `runId` query 参数。
- Trace response 增加 `runId`。
- 新增 `GET /api/chat/session/{sessionId}/runs`。
- Demo 脚本读取响应中的 `runId` 并查询精确 Trace。
- Trace UI 支持 `sessionId + runId` 最小查询。
验收:
@@ -448,8 +446,7 @@ tool_invocation(run_id, id)
- 传第一轮 `runId` 只返回第一轮 step/tool。
- 传第二轮 `runId` 只返回第二轮 step/tool。
- run 列表 API 只返回轻量 run 摘要,不展开 trace 明细。
- Demo 脚本输出摘要包含 `runId`。
- Trace UI 可以通过 URL 参数打开指定 run。
- Demo 脚本和 Trace UI 的 `runId` 最小适配按 OpenSpec tasks 放到 Phase 6,避免 Phase 3 同时混入前端/脚本范围。
### Phase 4:Feedback 和 CaseLibrary 绑定 run
@@ -475,7 +472,8 @@ tool_invocation(run_id, id)
任务:
- `AiOpsService` 生成并返回/透出 `runId`。
- AIOps `agent_step` / `tool_invocation` / rule evaluation 按 `runId` 隔离。
- AIOps `agent_step` / `tool_invocation` 按 `runId` 隔离。
- AIOps rule evaluation 写入当前 run 的 `diagnosis_run.self_evaluation.aiops_rule_evaluation`。
- AIOps Trace 查询兼容 `sessionId + runId`。
验收:
@@ -2,7 +2,7 @@
## sm-flow State
- Checkpoint: Discover
- Checkpoint: Apply / Phase 2
- Scale: complex
- Capability source: sm-flow built-in protocol for context/proposal; grill decisions are recorded from the confirmed user discussion in the issue thread.
- Change slug: `session-run-trace-isolation`
@@ -141,3 +141,13 @@ Audit conclusions:
- Clarified AIOps SSE compatibility: emit a metadata message containing `sessionId` and `runId` before report content while preserving the existing content stream shape.
- Clarified that run/session ownership is enforced by service-layer validation and indexed lookup in this change; database foreign keys are intentionally deferred to preserve compatibility with historical orphan detail rows and rollback.
- Clarified that `chat_session.expires_at` is nullable directory metadata / best-effort TTL snapshot, not mandatory persisted conversation history.
- Clarified document review findings before continuing Phase 2: OpenSpec task phases are authoritative over the older active issue phase sketch, and AIOps rule evaluation is stored in `diagnosis_run.self_evaluation.aiops_rule_evaluation`, not a separate table.
## Phase 2 Apply Notes
- Capability source: `openspec-apply-change` + sm-flow apply protocol. `codebase-retrieval` and LSP tools were not available in this session, so call-chain confirmation used OpenSpec context, `rg`, targeted file reads, compilation, focused tests, E2E, DB inspection, and logs.
- Implemented unified Chat execution context carrying `sessionId` and `runId` through `RunnableConfig.metadata` and `SessionContextHolder`.
- Switched valid Chat writes to create/update `chat_session` metadata and create one `diagnosis_run` per request.
- Switched Chat run completion, failure, self-evaluation, metrics, verifier support reads, gatekeeper validation, and evidence scoring to run-scoped data.
- Added `/api/chat` response `runId` and focused tests for valid run creation, invalid request no-run behavior, run-scoped trace consumers, and same-session multi-turn run creation.
- Phase 2 gate evidence is recorded in `phase-2-evidence.md`.
@@ -118,9 +118,12 @@ Compatibility:
5. Add indexes for `session_id`, `run_id`, latest-run lookup, and trace-detail lookup.
6. Deploy repository and read-path compatibility.
7. Switch Chat write path to `chat_session + diagnosis_run`.
8. Switch Trace, evaluation, feedback, case-library, demo, and UI paths.
9. Switch AIOps write path.
10. Verify no new rows are missing `run_id`; only then tighten application-level and, if safe, database-level non-null assumptions for new data.
8. Switch Chat evaluation and verifier support reads to run-scoped data as part of the Chat write-path cutover.
9. Switch Trace run resolution and run listing.
10. Switch Feedback and case-library paths.
11. Switch AIOps write path, including `diagnosis_run.self_evaluation.aiops_rule_evaluation`.
12. Switch demo scripts, Trace UI, and MVP docs.
13. Verify no new rows are missing `run_id`; only then tighten application-level and, if safe, database-level non-null assumptions for new data.
Rollback:
@@ -0,0 +1,129 @@
# Phase 2 Evidence: Chat Run Write Path
Date: 2026-07-10
## Scope
Phase 2 switched Chat writes from session-scoped execution state to run-scoped execution state:
- Chat creates/updates `chat_session` metadata.
- Each valid `/api/chat` execution creates one `diagnosis_run`.
- Chat execution context carries `sessionId + runId` through `RunnableConfig` and `SessionContextHolder`.
- `agent_step.run_id` and `tool_invocation.run_id` are written for Chat runs.
- Chat completion/failure/status/answer/self-evaluation/counts are written to `diagnosis_run`.
- verifier/gatekeeper/evaluation reads use run-scoped tool rows when `runId` is available.
- `/api/chat` response includes official `runId`.
## Static / Unit Verification
Commands:
```powershell
mvn -q clean test-compile
mvn -q "-Dtest=ChatControllerTest,ChatServiceSequentialAgentTest,ToolInvocationRecorderTest,ToolTraceSummaryServiceTest,ExecutorGatekeeperServiceTest" test
openspec validate session-run-trace-isolation --strict
```
Result:
- `test-compile` passed.
- Focused Phase 2 tests passed.
- OpenSpec strict validation passed.
Focused coverage:
- valid Chat creates `chat_session` and `diagnosis_run`;
- invalid blank Chat request returns before `ChatService`, so no run is created;
- same `sessionId` across two Chat turns creates two distinct `runId` values;
- `ToolInvocationRecorder` copies `runId` from execution context;
- verifier trace summary reads by run;
- gatekeeper validates by run;
- evaluation writes rule evaluation to `diagnosis_run.self_evaluation`.
## E2E Verification
Startup command:
```powershell
mvn spring-boot:run -Dspring-boot.run.profiles=mvp-demo
```
Process stdout/stderr:
- `target/e2e/phase2-mvn-20260710-185102.out.log`
- `target/e2e/phase2-mvn-20260710-185102.err.log`
Primary E2E session:
```text
sessionId = e2e-phase2-codex-20260710-1856
round 1 runId = run-24b6f04c-94a0-43cb-b94f-4f0141f9050d
round 2 runId = run-a7a2be77-697f-495a-a779-91afd1d8589c
```
HTTP evidence:
- `target/e2e/phase2-utf8-request-round1.json`
- `target/e2e/phase2-utf8-response-round1.json`
- `target/e2e/phase2-utf8-request-round2.json`
- `target/e2e/phase2-utf8-response-round2.json`
- both Chat responses returned `code=200`, `data.success=true`, the same `sessionId`, and distinct `runId` values.
Redis/session continuity evidence:
- `target/e2e/phase2-utf8-chat-session-response.json`
- response returned `messagePairCount=2`.
- logs show the second request entered with `会话历史消息对数: 1` and completed with `当前消息对数: 2`.
## Database Inspection
DB inspection used `scripts/query_mysql.py`.
Saved query outputs:
- `target/e2e/phase2-utf8-db-runs.txt`
- `target/e2e/phase2-utf8-db-chat-session.txt`
- `target/e2e/phase2-utf8-db-agent-steps.txt`
- `target/e2e/phase2-utf8-db-tool-invocations.txt`
- `target/e2e/phase2-utf8-db-missing-runid.txt`
Observed rows:
```text
diagnosis_run:
run-24b6f04c-94a0-43cb-b94f-4f0141f9050d | SUCCESS | CHAT | step_count=8 | tool_call_count=12
run-a7a2be77-697f-495a-a779-91afd1d8589c | SUCCESS | CHAT | step_count=2 | tool_call_count=1
chat_session:
e2e-phase2-codex-20260710-1856 | ACTIVE | message_pair_count=2
agent_step grouped by run_id:
run-24b6f04c-94a0-43cb-b94f-4f0141f9050d | 8
run-a7a2be77-697f-495a-a779-91afd1d8589c | 2
tool_invocation grouped by run_id:
run-24b6f04c-94a0-43cb-b94f-4f0141f9050d | 12
run-a7a2be77-697f-495a-a779-91afd1d8589c | 1
missing run_id for this session:
agent_step = 0
tool_invocation = 0
```
## Log Review
Saved log excerpts:
- `target/e2e/phase2-utf8-application-log-excerpt.txt`
- `target/e2e/phase2-utf8-mvn-log-excerpt.txt`
Findings:
- Chat logs show second turn reused Redis history for the same session.
- `EvaluationService` wrote scoring results to both run ids.
- No E2E-specific application exception was observed for `e2e-phase2-codex-20260710-1856`.
- Earlier `/actuator/health` probes produced expected 500/no-resource noise because the actuator health endpoint is not exposed; the E2E readiness check used `/api/chat` instead.
## Notes
`codebase-retrieval` and LSP tools were not available in this environment. Call-chain confirmation used OpenSpec context, `rg`, targeted file reads, compilation, focused tests, E2E, DB inspection, and logs.
@@ -26,7 +26,7 @@ Introduce a stable split between conversation state and execution state:
`runId` becomes an official API field:
- `/api/chat` returns `sessionId + runId`.
- `/api/ai_ops` emits or returns `runId` in the SSE-compatible protocol.
- `/api/ai_ops` emits an SSE-compatible metadata message containing `sessionId` and `runId` before report content.
- `GET /api/diagnosis/{sessionId}/trace` defaults to the latest run for compatibility.
- `GET /api/diagnosis/{sessionId}/trace?runId=run-...` returns the specified run after validating it belongs to the path `sessionId`.
- Feedback prefers `runId`; missing `runId` temporarily falls back to the latest run and returns both `fallbackToLatestRun=true` and the actual bound `runId`.
@@ -82,4 +82,3 @@ Level: L4 database/API contract migration with compatibility behavior.
- Evidence score, verifier inputs, and baseline metrics may change after run isolation because cross-round tool rows are no longer counted.
- Chat-only intermediate completion would leave AIOps as the remaining mixed-trace entry point; AIOps must be completed before overall archive.
- `case_library.diagnosis_id` becomes transitional: old data may contain `session_id`, new data contains `run_id`.
@@ -61,6 +61,7 @@ The system SHALL allow callers to query a diagnosis trace by `sessionId` alone f
- **WHEN** a caller requests `GET /api/diagnosis/{sessionId}/trace?runId=run-xxx`
- **THEN** the system SHALL validate that `runId` belongs to the path `sessionId`
- **AND** it SHALL return only the session summary, run summary, agent steps, tool invocations, self-evaluation, answer, and feedback for that run
- **AND** the session summary SHALL come from `chat_session` metadata when available, while the run summary SHALL come from `diagnosis_run`
#### Scenario: Trace rejects run from another session
- **WHEN** a caller requests a `runId` that belongs to a different `sessionId`
@@ -101,6 +102,7 @@ The system SHALL create and expose a diagnosis run for every valid `/api/ai_ops`
- **WHEN** `/api/ai_ops` starts a valid execution
- **THEN** the system SHALL create a `diagnosis_run` with `agent_flow=AI_OPS`
- **AND** AIOps agent steps, tool invocations, and rule evaluation SHALL be associated with that `run_id`
- **AND** AIOps rule evaluation SHALL be stored under `diagnosis_run.self_evaluation.aiops_rule_evaluation` for the current run
#### Scenario: AIOps SSE exposes runId
- **WHEN** `/api/ai_ops` streams response metadata to the caller
@@ -10,15 +10,15 @@
## 2. Chat Run Write Path
- [ ] 2.1 Add a unified execution context that carries both `sessionId` and `runId` through Chat service, Agent hooks, and tool recording.
- [ ] 2.2 Change valid `/api/chat` executions to create or update `chat_session` metadata and create one new `diagnosis_run`.
- [ ] 2.3 Change `AgentLoggingHook` to write `agent_step.run_id` for Chat runs while retaining `session_id`.
- [ ] 2.4 Change `ToolInvocationRecorder` and evidence tools to write `tool_invocation.run_id` for Chat runs while retaining `session_id`.
- [ ] 2.5 Change Chat completion, failure, answer, self-evaluation, duration, token, step, and tool count writes from `diagnosis_session` to the current `diagnosis_run`.
- [ ] 2.6 Change `ToolTraceSummaryService`, `ExecutorGatekeeperService`, and `EvaluationService` Chat reads from session-scoped tool rows to run-scoped tool rows.
- [ ] 2.7 Change `/api/chat` response DTO to include official `runId`.
- [ ] 2.8 Add focused tests for valid Chat run creation, invalid request no-run behavior, run-scoped counts, run-scoped verifier/gatekeeper/evaluation reads, and multi-turn context preservation.
- [ ] 2.9 Phase 2 gate: run focused tests plus a same-session two-round Chat E2E when needed, inspect DB with `scripts/query_mysql.py`, review `logs/`, update task status, archive phase evidence, and commit before starting Phase 3.
- [x] 2.1 Add a unified execution context that carries both `sessionId` and `runId` through Chat service, Agent hooks, and tool recording.
- [x] 2.2 Change valid `/api/chat` executions to create or update `chat_session` metadata and create one new `diagnosis_run`.
- [x] 2.3 Change `AgentLoggingHook` to write `agent_step.run_id` for Chat runs while retaining `session_id`.
- [x] 2.4 Change `ToolInvocationRecorder` and evidence tools to write `tool_invocation.run_id` for Chat runs while retaining `session_id`.
- [x] 2.5 Change Chat completion, failure, answer, self-evaluation, duration, token, step, and tool count writes from `diagnosis_session` to the current `diagnosis_run`.
- [x] 2.6 Change `ToolTraceSummaryService`, `ExecutorGatekeeperService`, and `EvaluationService` Chat reads from session-scoped tool rows to run-scoped tool rows.
- [x] 2.7 Change `/api/chat` response DTO to include official `runId`.
- [x] 2.8 Add focused tests for valid Chat run creation, invalid request no-run behavior, run-scoped counts, run-scoped verifier/gatekeeper/evaluation reads, and multi-turn context preservation.
- [x] 2.9 Phase 2 gate: run focused tests plus a same-session two-round Chat E2E when needed, inspect DB with `scripts/query_mysql.py`, review `logs/`, update task status, archive phase evidence, and commit before starting Phase 3.
## 3. Trace Read Path and Run Listing
@@ -45,7 +45,7 @@
- [ ] 5.1 Change valid `/api/ai_ops` executions to create `diagnosis_run` with `agent_flow=AI_OPS`.
- [ ] 5.2 Expose `runId` in the AIOps SSE-compatible metadata stream while preserving existing report streaming.
- [ ] 5.3 Propagate `runId` through AIOps Agent hooks and tool recording.
- [ ] 5.4 Change AIOps final report, status, counts, and `aiops_rule_evaluation` writes to the current `diagnosis_run`.
- [ ] 5.4 Change AIOps final report, status, counts, and `diagnosis_run.self_evaluation.aiops_rule_evaluation` writes to the current run.
- [ ] 5.5 Add tests for repeated AIOps executions with the same `sessionId` and run-scoped rule evaluation.
- [ ] 5.6 Phase 5 gate: run focused AIOps tests and E2E when needed, inspect DB/logs, update task status, archive phase evidence, and commit before starting Phase 6.
@@ -96,10 +96,11 @@ public class ChatController {
// 更新会话历史
session.addChatMessagePair(request.getQuestion(), fullAnswer, MAX_WINDOW_SIZE);
sessionManager.updateSession(session);
chatService.syncChatSessionMetadata(session.getSessionId(), session.getMessagePairCount());
logger.info("已更新会话历史 - SessionId: {}, 当前消息对数: {}",
session.getSessionId(), session.getMessagePairCount());
return ResponseEntity.ok(ApiResponse.success(ChatResponse.success(fullAnswer, result.sessionId())));
return ResponseEntity.ok(ApiResponse.success(ChatResponse.success(fullAnswer, result.sessionId(), result.runId())));
} catch (Exception e) {
logger.error("对话失败", e);
@@ -413,12 +414,14 @@ public class ChatController {
private String answer;
private String errorMessage;
private String sessionId;
private String runId;
public static ChatResponse success(String answer, String sessionId) {
public static ChatResponse success(String answer, String sessionId, String runId) {
ChatResponse response = new ChatResponse();
response.setSuccess(true);
response.setAnswer(answer);
response.setSessionId(sessionId);
response.setRunId(runId);
return response;
}
@@ -47,11 +47,13 @@ public class AgentLoggingHook extends MessagesModelHook {
@Override
public AgentCommand beforeModel(List<Message> previousMessages, RunnableConfig config) {
String sessionId = resolveSessionId(config);
String runId = resolveRunId(config);
String traceScopeId = traceScopeId(sessionId, runId);
boolean hasSession = sessionId != null;
int stepIndex = 0;
if (hasSession) {
stepIndex = stepCounters.merge(sessionId, 0, (oldValue, ignored) -> oldValue + 1);
if (traceScopeId != null) {
stepIndex = stepCounters.merge(traceScopeId, 0, (oldValue, ignored) -> oldValue + 1);
}
log.info("========================================");
@@ -73,12 +75,13 @@ public class AgentLoggingHook extends MessagesModelHook {
try {
AgentStep step = AgentStep.builder()
.sessionId(sessionId)
.runId(runId)
.stepIndex(stepIndex)
.agentName(agentName)
.modelInput(buildModelInputSummary(previousMessages))
.build();
AgentStep saved = agentStepRepository.save(step);
pendingSteps.put(sessionId + "_" + stepIndex, Map.of(
pendingSteps.put(stepKey(traceScopeId, stepIndex), Map.of(
"stepId", saved.getId(),
"startTime", System.currentTimeMillis()
));
@@ -93,7 +96,9 @@ public class AgentLoggingHook extends MessagesModelHook {
@Override
public AgentCommand afterModel(List<Message> previousMessages, RunnableConfig config) {
String sessionId = resolveSessionId(config);
int stepIndex = sessionId == null ? 0 : stepCounters.getOrDefault(sessionId, 0);
String runId = resolveRunId(config);
String traceScopeId = traceScopeId(sessionId, runId);
int stepIndex = traceScopeId == null ? 0 : stepCounters.getOrDefault(traceScopeId, 0);
log.info("========================================");
log.info("*** [AgentTrace] agent={}, phase=after_model, stepIndex={}", agentName, stepIndex);
@@ -122,7 +127,7 @@ public class AgentLoggingHook extends MessagesModelHook {
log.info("========================================");
if (sessionId != null) {
String stepKey = sessionId + "_" + stepIndex;
String stepKey = stepKey(traceScopeId, stepIndex);
Map<String, Object> pending = pendingSteps.remove(stepKey);
if (pending != null) {
try {
@@ -162,6 +167,23 @@ public class AgentLoggingHook extends MessagesModelHook {
.orElseGet(SessionContextHolder::getSessionId);
}
private String resolveRunId(RunnableConfig config) {
return config.metadata("runId")
.map(Object::toString)
.orElseGet(SessionContextHolder::getRunId);
}
private String traceScopeId(String sessionId, String runId) {
if (runId != null && !runId.isBlank()) {
return runId;
}
return sessionId;
}
private String stepKey(String traceScopeId, int stepIndex) {
return traceScopeId + "_" + stepIndex;
}
private AssistantMessage findLastAssistant(List<Message> previousMessages) {
for (int i = previousMessages.size() - 1; i >= 0; i--) {
if (previousMessages.get(i) instanceof AssistantMessage assistantMessage) {
@@ -59,13 +59,17 @@ public class VerifierInputHook extends MessagesModelHook {
String sessionId = config.metadata("sessionId")
.map(Object::toString)
.orElseGet(SessionContextHolder::getSessionId);
String runId = config.metadata("runId")
.map(Object::toString)
.orElseGet(SessionContextHolder::getRunId);
String executorFinalAnswer = VerifierContextHolder.getExecutorFinalAnswer();
if (executorFinalAnswer == null || executorFinalAnswer.isBlank()) {
executorFinalAnswer = extractLastAssistantText(previousMessages);
}
List<Map<String, Object>> toolTraceSummary =
toolTraceSummaryService.buildVerifierTraceSummary(sessionId, executorFinalAnswer);
List<Map<String, Object>> toolTraceSummary = runId == null || runId.isBlank()
? toolTraceSummaryService.buildVerifierTraceSummary(sessionId, executorFinalAnswer)
: toolTraceSummaryService.buildVerifierTraceSummaryForRun(runId, executorFinalAnswer);
VerifierContextHolder.setToolTraceSummary(toolTraceSummary);
ExecutorOutputParseResult parseResult = parseExecutorOutput(executorFinalAnswer);
@@ -76,7 +80,7 @@ public class VerifierInputHook extends MessagesModelHook {
VerifierContextHolder.setExecutorStructuredOutput(parseResult.structuredOutput());
VerifierContextHolder.setExecutorOutputParseStatus(parseResult.status());
Map<String, Object> gatekeeperResult = runGatekeeper(sessionId, parseResult);
Map<String, Object> gatekeeperResult = runGatekeeper(sessionId, runId, parseResult);
VerifierContextHolder.setGatekeeperResult(gatekeeperResult);
Map<String, Object> verifierInput = new LinkedHashMap<>();
@@ -96,11 +100,14 @@ public class VerifierInputHook extends MessagesModelHook {
}
}
private Map<String, Object> runGatekeeper(String sessionId, ExecutorOutputParseResult parseResult) {
private Map<String, Object> runGatekeeper(String sessionId, String runId, ExecutorOutputParseResult parseResult) {
if (executorGatekeeperService == null) {
return passGatekeeperResult();
}
try {
if (runId != null && !runId.isBlank()) {
return executorGatekeeperService.validateRun(runId, parseResult.structuredOutput(), parseResult.status());
}
return executorGatekeeperService.validate(sessionId, parseResult.structuredOutput(), parseResult.status());
} catch (Exception e) {
log.error("Gatekeeper validation failed unexpectedly", e);
@@ -14,14 +14,16 @@ import com.superbiz.agent.agent.tool.DateTimeTools;
import com.superbiz.agent.agent.tool.InternalDocsTools;
import com.superbiz.agent.agent.tool.QueryLogsTools;
import com.superbiz.agent.agent.tool.QueryMetricsTools;
import com.superbiz.agent.domain.entity.DiagnosisSession;
import com.superbiz.agent.domain.entity.ChatSession;
import com.superbiz.agent.domain.entity.DiagnosisRun;
import com.superbiz.agent.hook.AgentLoggingHook;
import com.superbiz.agent.hook.PlannerSkillMetadataHook;
import com.superbiz.agent.hook.TokenTrackingChatModel;
import com.superbiz.agent.hook.TokenUsageHolder;
import com.superbiz.agent.hook.VerifierInputHook;
import com.superbiz.agent.repository.AgentStepRepository;
import com.superbiz.agent.repository.DiagnosisSessionRepository;
import com.superbiz.agent.repository.ChatSessionRepository;
import com.superbiz.agent.repository.DiagnosisRunRepository;
import com.superbiz.agent.repository.ToolInvocationRepository;
import com.superbiz.agent.tool.LookupKnowledgeTool;
import com.superbiz.agent.tool.RetrievedDocTracker;
@@ -43,6 +45,7 @@ import org.springframework.stereotype.Service;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.time.LocalDateTime;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
@@ -62,8 +65,8 @@ public class ChatService {
private static final String DEGRADED_PREFIX = "当前无法基于已获取证据生成可靠结论,建议人工介入。";
private static final String CHAT_PROMPT_AUDIT_VERSION = "chat-prompts-v1";
/** 封装 answer + 后端生成的 sessionId,用于 feedback 关联 */
public record ChatResult(String answer, String sessionId) {}
/** 封装 answer + 后端生成的 sessionId/runId,用于 feedback 关联 */
public record ChatResult(String answer, String sessionId, String runId) {}
@Autowired
private InternalDocsTools internalDocsTools;
@@ -74,7 +77,7 @@ public class ChatService {
@Autowired
private QueryMetricsTools queryMetricsTools;
@Autowired(required = false) // Mock 模式下才注册,所以设置为 optional,真实环境通过mcp配置注入
@Autowired(required = false) // Mock 模式下才注册,所以设置为 optional,真实环境通过 mcp 配置注入
private QueryLogsTools queryLogsTools;
@Autowired(required = false)
@@ -87,7 +90,10 @@ public class ChatService {
private LookupKnowledgeTool lookupKnowledgeTool;
@Autowired
private DiagnosisSessionRepository diagnosisSessionRepository;
private ChatSessionRepository chatSessionRepository;
@Autowired
private DiagnosisRunRepository diagnosisRunRepository;
@Autowired
private AgentStepRepository agentStepRepository;
@@ -131,7 +137,7 @@ public class ChatService {
@PostConstruct
public void init() {
// 加载 Prompt
// 鍔犺浇 Prompt
try {
chatPlannerPrompt = new String(
new ClassPathResource("prompts/chat-planner-prompt.md").getInputStream().readAllBytes(),
@@ -186,7 +192,7 @@ public class ChatService {
String role = msg.get("role");
String content = msg.get("content");
// 🔧 过滤时间查询相关的历史消息,避免 LLM 复用旧的时间信息
// 过滤时间查询相关的历史消息,避免 LLM 复用旧的时间信息
if ("user".equals(role) && isTimeQuery(content)) {
continue; // 跳过时间查询问题
}
@@ -229,7 +235,7 @@ public class ChatService {
return false;
}
// 匹配日期时间格式:2026年5月31日、15:57、下午3点 等
return content.matches(".*(\\d{4}年\\d{1,2}月\\d{1,2}日|\\d{1,2}:\\d{2}|[上下午]+\\d{1,2}[点时]).*");
return content.matches(".*(\\d{4}.*\\d{1,2}.*\\d{1,2}.*|\\d{1,2}:\\d{2}|[上下]午\\s*\\d{1,2}[点时]).*");
}
/**
@@ -250,7 +256,7 @@ public class ChatService {
}
/**
* 获取工具回调列表,mcp服务提供的工具
* 获取工具回调列表,mcp 服务提供的工具
*/
public ToolCallback[] getToolCallbacks() {
if (tools == null) {
@@ -260,7 +266,7 @@ public class ChatService {
}
/**
* 记录可用工具列表:mcp服务提供的工具
* 记录可用工具列表:mcp 服务提供的工具
*/
public void logAvailableTools() {
if (tools == null) {
@@ -295,7 +301,7 @@ public class ChatService {
* 执行 ReactAgent 对话(非流式)
* @param agent ReactAgent 实例
* @param question 用户问题
* @return ChatResult(answer + sessionId)
* @return ChatResult(answer + sessionId + runId)
*/
public ChatResult executeChat(ReactAgent agent, String question) throws GraphRunnerException {
return executeChat(agent, question, null);
@@ -303,22 +309,24 @@ public class ChatService {
public ChatResult executeChat(ReactAgent agent, String question, String requestedSessionId) throws GraphRunnerException {
logger.info("========================================");
logger.info("📝 用户问题: {}", question);
logger.info("用户问题: {}", question);
String sessionId = resolveSessionId(requestedSessionId);
String runId = newRunId();
long startTime = System.currentTimeMillis();
// 创建或更新诊断会话
DiagnosisSession session = startDiagnosisSession(sessionId, question);
diagnosisSessionRepository.save(session);
// 创建或更新 Chat Session 元数据,并创建本次诊断 run
ensureChatSession(sessionId, null);
DiagnosisRun run = startDiagnosisRun(sessionId, runId, question);
// 设置 ThreadLocal 上下文(LookupKnowledgeTool 通过此获取 sessionId)
SessionContextHolder.setSessionId(sessionId);
// 设置 ThreadLocal 上下文,供工具和 Hook 读取 sessionId/runId
SessionContextHolder.setContext(sessionId, runId);
try {
// 通过 RunnableConfig 将 sessionId 传入 Hook(线程安全,异步也兼容)
// 通过 RunnableConfig 将 sessionId/runId 传入 Hook
var config = RunnableConfig.builder()
.addMetadata("sessionId", sessionId)
.addMetadata("runId", runId)
.build();
var response = agent.call(question, config);
@@ -326,23 +334,27 @@ public class ChatService {
String answer = response.getText();
// 更新诊断会话
session.setStatus("SUCCESS");
session.setAnswer(answer);
session.setTotalDurationMs((int) duration);
backfillSessionMetrics(session);
diagnosisSessionRepository.save(session);
// 更新诊断 run
run.setStatus("SUCCESS");
run.setAnswer(answer);
run.setTotalDurationMs((int) duration);
backfillRunMetrics(run);
diagnosisRunRepository.save(run);
evaluationService.evaluate(sessionId, answer);
evaluationService.evaluateRun(runId, answer);
logger.info("⏱️ 总耗时: {} ms", duration);
logger.info("📏 输出长度: {} 字符", answer.length());
logger.info("总耗时: {} ms", duration);
logger.info("输出长度: {} 字符", answer.length());
logger.info("========================================");
return new ChatResult(answer, sessionId);
return new ChatResult(answer, sessionId, runId);
} catch (Exception e) {
session.setStatus("FAILED");
diagnosisSessionRepository.save(session);
String errorAnswer = "Execution failed: " + e.getMessage();
run.setStatus("FAILED");
run.setAnswer(errorAnswer);
run.setTotalDurationMs((int) (System.currentTimeMillis() - startTime));
backfillRunMetrics(run);
diagnosisRunRepository.save(run);
throw e;
} finally {
retrievedDocTracker.clearSession(sessionId);
@@ -356,7 +368,7 @@ public class ChatService {
* @param toolCallbacks 工具回调
* @param question 用户问题
* @param history 历史消息
* @return ChatResult(answer + sessionId)
* @return ChatResult(answer + sessionId + runId)
*/
public ChatResult executeChatWithStrategy(ChatModel chatModel, ToolCallback[] toolCallbacks,
String question, List<Map<String, String>> history) throws GraphRunnerException {
@@ -367,10 +379,10 @@ public class ChatService {
String question, List<Map<String, String>> history,
String requestedSessionId) throws GraphRunnerException {
if (QuestionComplexity.isComplex(question)) {
logger.info("📊 问题判定为复杂,使用多 Agent(Planner + Executor)执行");
logger.info("问题判定为复杂,使用多 Agent(Planner + Executor)执行");
return executeChatComplex(chatModel, toolCallbacks, question, history, requestedSessionId);
} else {
logger.info("📊 问题判定为简单,使用单 Agent 执行");
logger.info("问题判定为简单,使用单 Agent 执行");
String systemPrompt = buildSystemPrompt(history);
ReactAgent agent = createReactAgent(chatModel, systemPrompt);
return executeChat(agent, question, requestedSessionId);
@@ -389,12 +401,13 @@ public class ChatService {
String question, List<Map<String, String>> history,
String requestedSessionId) throws GraphRunnerException {
String sessionId = resolveSessionId(requestedSessionId);
String runId = newRunId();
long startTime = System.currentTimeMillis();
DiagnosisSession session = startDiagnosisSession(sessionId, question);
diagnosisSessionRepository.save(session);
ensureChatSession(sessionId, history == null ? null : history.size() / 2);
DiagnosisRun run = startDiagnosisRun(sessionId, runId, question);
SessionContextHolder.setSessionId(sessionId);
SessionContextHolder.setContext(sessionId, runId);
VerifierContextHolder.setOriginalQuery(question);
VerifierContextHolder.setRetryContext(null);
VerifierContextHolder.setExecutorFinalAnswer(null);
@@ -405,6 +418,7 @@ public class ChatService {
String answer = null;
RunnableConfig config = RunnableConfig.builder()
.addMetadata("sessionId", sessionId)
.addMetadata("runId", runId)
.build();
for (int round = 1; round <= 2; round++) {
@@ -427,7 +441,7 @@ public class ChatService {
finalDecision = buildVerifierFallbackDecision(round, "workflow 未返回有效状态");
ComposerRenderResult renderResult = buildFixedFallbackAnswer(question, finalDecision);
answer = renderResult.answer();
persistVerifierEvaluation(session, finalDecision, round, renderResult.audit());
persistVerifierEvaluation(run, finalDecision, round, renderResult.audit());
break;
}
@@ -449,21 +463,21 @@ public class ChatService {
finalDecision = buildVerifierFallbackDecision(round, "verifier_output 缺失或无法解析");
ComposerRenderResult renderResult = buildFixedFallbackAnswer(question, finalDecision);
answer = renderResult.answer();
persistVerifierEvaluation(session, finalDecision, round, renderResult.audit());
persistVerifierEvaluation(run, finalDecision, round, renderResult.audit());
break;
}
if ("PASS".equals(finalDecision.verdict())) {
ComposerRenderResult renderResult = composeFinalAnswer(chatModel, question, finalDecision, config);
answer = renderResult.answer();
persistVerifierEvaluation(session, finalDecision, round, renderResult.audit());
persistVerifierEvaluation(run, finalDecision, round, renderResult.audit());
break;
}
if ("REJECT".equals(finalDecision.verdict())) {
ComposerRenderResult renderResult = composeFinalAnswer(chatModel, question, finalDecision, config);
answer = renderResult.answer();
persistVerifierEvaluation(session, finalDecision, round, renderResult.audit());
persistVerifierEvaluation(run, finalDecision, round, renderResult.audit());
break;
}
@@ -473,12 +487,12 @@ public class ChatService {
if (!shouldRetry) {
ComposerRenderResult renderResult = composeFinalAnswer(chatModel, question, finalDecision, config);
answer = renderResult.answer();
persistVerifierEvaluation(session, finalDecision, round, renderResult.audit());
persistVerifierEvaluation(run, finalDecision, round, renderResult.audit());
break;
}
retryContext = buildRetryContext(finalDecision);
persistVerifierEvaluation(session, finalDecision, round);
persistVerifierEvaluation(run, finalDecision, round);
}
long duration = System.currentTimeMillis() - startTime;
@@ -487,24 +501,28 @@ public class ChatService {
answer = "抱歉,多 Agent 分析未能生成有效结论。";
}
session.setStatus("SUCCESS");
session.setAnswer(answer);
session.setTotalDurationMs((int) duration);
backfillSessionMetrics(session);
diagnosisSessionRepository.save(session);
run.setStatus("SUCCESS");
run.setAnswer(answer);
run.setTotalDurationMs((int) duration);
backfillRunMetrics(run);
diagnosisRunRepository.save(run);
evaluationService.evaluate(sessionId, answer);
evaluationService.evaluateRun(runId, answer);
logger.info("⏱️ 多 Agent 总耗时: {} ms", duration);
logger.info("📏 输出长度: {} 字符", answer.length());
logger.info("多 Agent 总耗时: {} ms", duration);
logger.info("输出长度: {} 字符", answer.length());
return new ChatResult(answer, sessionId);
return new ChatResult(answer, sessionId, runId);
} catch (Exception e) {
session.setStatus("FAILED");
diagnosisSessionRepository.save(session);
String errorAnswer = "Execution failed: " + e.getMessage();
run.setStatus("FAILED");
run.setAnswer(errorAnswer);
run.setTotalDurationMs((int) (System.currentTimeMillis() - startTime));
backfillRunMetrics(run);
diagnosisRunRepository.save(run);
logger.error("多 Agent 执行失败", e);
return new ChatResult("执行失败: " + e.getMessage(), sessionId);
return new ChatResult(errorAnswer, sessionId, runId);
} finally {
retrievedDocTracker.clearSession(sessionId);
SessionContextHolder.clear();
@@ -516,7 +534,7 @@ public class ChatService {
String retryContext) {
StringBuilder prompt = new StringBuilder(chatPlannerPrompt);
// 注入 knowledge map
// 娉ㄥ叆 knowledge map
String knowledgeMap = knowledgeDomainService.buildKnowledgeMap();
if (!knowledgeMap.isBlank()) {
prompt.append("\n\n## 可用知识库\n\n").append(knowledgeMap);
@@ -530,7 +548,7 @@ public class ChatService {
prompt.append("--- 对话历史结束 ---\n");
}
if (retryContext != null && !retryContext.isBlank()) {
prompt.append("\n\n--- 本轮补证据约束 ---\n").append(retryContext).append("\n");
prompt.append("\n\n--- 鏈疆琛ヨ瘉鎹害鏉?---\n").append(retryContext).append("\n");
}
return ReactAgent.builder()
.name("chat_planner")
@@ -580,7 +598,7 @@ public class ChatService {
prompt.append("--- 对话历史结束 ---\n");
}
if (retryContext != null && !retryContext.isBlank()) {
prompt.append("\n\n--- 本轮补证据约束 ---\n").append(retryContext).append("\n");
prompt.append("\n\n--- 鏈疆琛ヨ瘉鎹害鏉?---\n").append(retryContext).append("\n");
}
return ReactAgent.builder()
.name("chat_executor")
@@ -620,21 +638,40 @@ public class ChatService {
return UUID.randomUUID().toString().substring(0, 8);
}
private DiagnosisSession startDiagnosisSession(String sessionId, String question) {
DiagnosisSession session = diagnosisSessionRepository.findBySessionId(sessionId)
.orElseGet(() -> DiagnosisSession.builder()
private String newRunId() {
return "run-" + UUID.randomUUID();
}
public void syncChatSessionMetadata(String sessionId, Integer messagePairCount) {
if (sessionId == null || sessionId.isBlank()) {
return;
}
ensureChatSession(sessionId, messagePairCount);
}
private ChatSession ensureChatSession(String sessionId, Integer messagePairCount) {
ChatSession chatSession = chatSessionRepository.findBySessionId(sessionId)
.orElseGet(() -> ChatSession.builder()
.sessionId(sessionId)
.agentFlow("CHAT")
.status("ACTIVE")
.build());
session.setQuery(question);
session.setStatus("RUNNING");
session.setAgentFlow("CHAT");
session.setAnswer(null);
session.setTotalDurationMs(null);
session.setTotalTokenCount(null);
session.setStepCount(null);
session.setToolCallCount(null);
return session;
chatSession.setStatus("ACTIVE");
chatSession.setLastActiveAt(LocalDateTime.now());
if (messagePairCount != null) {
chatSession.setMessagePairCount(messagePairCount);
}
return chatSessionRepository.save(chatSession);
}
private DiagnosisRun startDiagnosisRun(String sessionId, String runId, String question) {
DiagnosisRun run = DiagnosisRun.builder()
.runId(runId)
.sessionId(sessionId)
.query(question)
.status("RUNNING")
.agentFlow("CHAT")
.build();
return diagnosisRunRepository.save(run);
}
private String buildWorkflowInput(String question, String retryContext) {
@@ -867,11 +904,11 @@ public class ChatService {
.orElse(null);
}
private void persistVerifierEvaluation(DiagnosisSession session, VerifierDecision decision, int round) {
persistVerifierEvaluation(session, decision, round, null);
private void persistVerifierEvaluation(DiagnosisRun run, VerifierDecision decision, int round) {
persistVerifierEvaluation(run, decision, round, null);
}
private void persistVerifierEvaluation(DiagnosisSession session, VerifierDecision decision, int round,
private void persistVerifierEvaluation(DiagnosisRun run, VerifierDecision decision, int round,
Map<String, Object> composerOutput) {
if (decision == null) {
return;
@@ -899,9 +936,9 @@ public class ChatService {
verifierEvaluation.put("composer_output", composerOutput);
}
String merged = selfEvaluationMergeService.mergeVerifierEvaluation(session.getSelfEvaluation(), verifierEvaluation);
session.setSelfEvaluation(merged);
diagnosisSessionRepository.save(session);
String merged = selfEvaluationMergeService.mergeVerifierEvaluation(run.getSelfEvaluation(), verifierEvaluation);
run.setSelfEvaluation(merged);
diagnosisRunRepository.save(run);
}
private Map<String, Object> promptAuditSnapshot() {
@@ -1244,7 +1281,10 @@ public class ChatService {
private List<String> buildNextStepSuggestionsFromTrace() {
List<String> suggestions = new ArrayList<>();
List<Map<String, Object>> toolSummary = toolTraceSummaryService.buildVerifierTraceSummary(SessionContextHolder.getSessionId(), null);
String runId = SessionContextHolder.getRunId();
List<Map<String, Object>> toolSummary = runId == null || runId.isBlank()
? toolTraceSummaryService.buildVerifierTraceSummary(SessionContextHolder.getSessionId(), null)
: toolTraceSummaryService.buildVerifierTraceSummaryForRun(runId, null);
boolean hasKnowledgeTool = toolSummary.stream().anyMatch(item -> "lookup_knowledge".equals(item.get("tool_name")));
boolean hasFailedEvidence = toolSummary.stream().anyMatch(item -> !Boolean.TRUE.equals(item.get("success")));
@@ -1309,11 +1349,11 @@ public class ChatService {
private record ComposerRenderResult(String answer, Map<String, Object> audit) {
}
/** 从 agent_step 和 tool_invocation 汇总指标回填 diagnosis_session */
private void backfillSessionMetrics(DiagnosisSession session) {
/** 从 agent_step 和 tool_invocation 汇总指标回填 diagnosis_run */
private void backfillRunMetrics(DiagnosisRun run) {
try {
List<com.superbiz.agent.domain.entity.AgentStep> steps =
agentStepRepository.findBySessionIdOrderByStepIndex(session.getSessionId());
agentStepRepository.findByRunIdOrderByStepIndex(run.getRunId());
int totalTokens = 0;
int stepCount = 0;
@@ -1321,12 +1361,12 @@ public class ChatService {
stepCount++;
if (s.getTokenCount() != null) totalTokens += s.getTokenCount();
}
long toolCallCount = toolInvocationRepository.countBySessionId(session.getSessionId());
session.setTotalTokenCount(totalTokens);
session.setStepCount(stepCount);
session.setToolCallCount(Math.toIntExact(toolCallCount));
long toolCallCount = toolInvocationRepository.countByRunId(run.getRunId());
run.setTotalTokenCount(totalTokens);
run.setStepCount(stepCount);
run.setToolCallCount(Math.toIntExact(toolCallCount));
} catch (Exception e) {
logger.warn("回填会话指标失败: sessionId={}", session.getSessionId(), e);
logger.warn("Failed to backfill run metrics: runId={}", run.getRunId(), e);
}
}
}
@@ -1,7 +1,9 @@
package com.superbiz.agent.service;
import com.superbiz.agent.domain.entity.DiagnosisSession;
import com.superbiz.agent.domain.entity.DiagnosisRun;
import com.superbiz.agent.domain.entity.ToolInvocation;
import com.superbiz.agent.repository.DiagnosisRunRepository;
import com.superbiz.agent.repository.DiagnosisSessionRepository;
import com.superbiz.agent.repository.ToolInvocationRepository;
import org.slf4j.Logger;
@@ -29,6 +31,9 @@ public class EvaluationService {
@Autowired
private DiagnosisSessionRepository diagnosisSessionRepository;
@Autowired
private DiagnosisRunRepository diagnosisRunRepository;
@Autowired
private ToolInvocationRepository toolInvocationRepository;
@@ -40,7 +45,7 @@ public class EvaluationService {
diagnosisSessionRepository.findBySessionId(sessionId).ifPresent(session -> {
try {
List<ToolInvocation> toolInvocations = toolInvocationRepository.findBySessionId(sessionId);
Map<String, Object> ruleEvaluation = evaluateWithRules(session, toolInvocations);
Map<String, Object> ruleEvaluation = evaluateWithRules(session.getStatus(), toolInvocations);
String merged = selfEvaluationMergeService.mergeRuleEvaluation(session.getSelfEvaluation(), ruleEvaluation);
session.setSelfEvaluation(merged);
diagnosisSessionRepository.save(session);
@@ -51,14 +56,30 @@ public class EvaluationService {
});
}
@Async
public void evaluateRun(String runId, String answer) {
diagnosisRunRepository.findByRunId(runId).ifPresent(run -> {
try {
List<ToolInvocation> toolInvocations = toolInvocationRepository.findByRunIdOrderByIdAsc(runId);
Map<String, Object> ruleEvaluation = evaluateWithRules(run.getStatus(), toolInvocations);
String merged = selfEvaluationMergeService.mergeRuleEvaluation(run.getSelfEvaluation(), ruleEvaluation);
run.setSelfEvaluation(merged);
diagnosisRunRepository.save(run);
logger.info("证据评分已写入: runId={}, result={}", runId, merged);
} catch (Exception e) {
logger.error("评分失败: runId={}", runId, e);
}
});
}
// -------------------------------------------------------------------------
// 规则引擎(事实层)
// -------------------------------------------------------------------------
private Map<String, Object> evaluateWithRules(DiagnosisSession session, List<ToolInvocation> invocations) {
private Map<String, Object> evaluateWithRules(String status, List<ToolInvocation> invocations) {
List<Map<String, Object>> factors = new ArrayList<>();
if ("FAILED".equals(session.getStatus())) {
if ("FAILED".equals(status)) {
factors.add(factor("execution_failed", -100, "执行失败"));
return buildResult(0, factors);
}
@@ -61,7 +61,25 @@ public class ExecutorGatekeeperService {
GatekeeperResult result = new GatekeeperResult(ruleCatalog);
validateSchema(structuredOutput, parseStatus, result);
if (structuredOutput != null) {
validateInvocationRefs(sessionId, structuredOutput, result);
List<ToolInvocation> invocations = sessionId == null || sessionId.isBlank()
? List.of()
: toolInvocationRepository.findBySessionIdOrderByIdAsc(sessionId);
validateInvocationRefs("session_id", sessionId, invocations, structuredOutput, result);
importWarnings(structuredOutput, result);
}
return result.toMap();
}
public Map<String, Object> validateRun(String runId,
Map<String, Object> structuredOutput,
Map<String, Object> parseStatus) {
GatekeeperResult result = new GatekeeperResult(ruleCatalog);
validateSchema(structuredOutput, parseStatus, result);
if (structuredOutput != null) {
List<ToolInvocation> invocations = runId == null || runId.isBlank()
? List.of()
: toolInvocationRepository.findByRunIdOrderByIdAsc(runId);
validateInvocationRefs("run_id", runId, invocations, structuredOutput, result);
importWarnings(structuredOutput, result);
}
return result.toMap();
@@ -134,15 +152,18 @@ public class ExecutorGatekeeperService {
requireArray(structuredOutput, "missing_info", result);
}
private void validateInvocationRefs(String sessionId, Map<String, Object> structuredOutput, GatekeeperResult result) {
if (sessionId == null || sessionId.isBlank()) {
result.fail(RULE_INVOCATION_REF, "session_id", "session id is required to validate source_invocation_id",
private void validateInvocationRefs(String scopeName,
String scopeId,
List<ToolInvocation> invocations,
Map<String, Object> structuredOutput,
GatekeeperResult result) {
if (scopeId == null || scopeId.isBlank()) {
result.fail(RULE_INVOCATION_REF, scopeName, scopeName + " is required to validate source_invocation_id",
SEVERITY_LOW_CONFID);
return;
}
Map<Long, ToolInvocation> validInvocations = toolInvocationRepository.findBySessionIdOrderByIdAsc(sessionId)
.stream()
Map<Long, ToolInvocation> validInvocations = invocations.stream()
.filter(invocation -> invocation.getId() != null)
.collect(Collectors.toMap(ToolInvocation::getId, Function.identity(), (left, right) -> left));
Object claimsValue = structuredOutput.get("claims");
@@ -48,6 +48,9 @@ public class ToolInvocationRecorder {
if (invocation.getSessionId() == null || invocation.getSessionId().isBlank()) {
invocation.setSessionId(SessionContextHolder.getSessionId());
}
if (invocation.getRunId() == null || invocation.getRunId().isBlank()) {
invocation.setRunId(SessionContextHolder.getRunId());
}
if (invocation.getSessionId() == null || invocation.getSessionId().isBlank()) {
log.debug("Skip tool_invocation without sessionId: tool={}", invocation.getToolName());
return;
@@ -42,6 +42,19 @@ public class ToolTraceSummaryService {
}
List<ToolInvocation> invocations = toolInvocationRepository.findBySessionIdOrderByIdAsc(sessionId);
return buildVerifierTraceSummary(invocations, executorFinalAnswer);
}
public List<Map<String, Object>> buildVerifierTraceSummaryForRun(String runId, String executorFinalAnswer) {
if (runId == null || runId.isBlank()) {
return List.of();
}
List<ToolInvocation> invocations = toolInvocationRepository.findByRunIdOrderByIdAsc(runId);
return buildVerifierTraceSummary(invocations, executorFinalAnswer);
}
private List<Map<String, Object>> buildVerifierTraceSummary(List<ToolInvocation> invocations, String executorFinalAnswer) {
if (invocations.isEmpty()) {
return List.of();
}
@@ -14,16 +14,31 @@ package com.superbiz.agent.util;
public class SessionContextHolder {
private static final ThreadLocal<String> SESSION_ID = new ThreadLocal<>();
private static final ThreadLocal<String> RUN_ID = new ThreadLocal<>();
public static void setContext(String sessionId, String runId) {
setSessionId(sessionId);
setRunId(runId);
}
public static void setSessionId(String sessionId) {
SESSION_ID.set(sessionId);
}
public static void setRunId(String runId) {
RUN_ID.set(runId);
}
public static String getSessionId() {
return SESSION_ID.get();
}
public static String getRunId() {
return RUN_ID.get();
}
public static void clear() {
SESSION_ID.remove();
RUN_ID.remove();
}
}
@@ -0,0 +1,32 @@
package com.superbiz.agent.controller;
import com.superbiz.agent.service.ChatService;
import org.junit.jupiter.api.Test;
import org.springframework.http.ResponseEntity;
import org.springframework.test.util.ReflectionTestUtils;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verifyNoInteractions;
class ChatControllerTest {
@Test
void blankChatRequestReturnsErrorBeforeCreatingRun() {
ChatController controller = new ChatController();
ChatService chatService = mock(ChatService.class);
ReflectionTestUtils.setField(controller, "chatService", chatService);
ChatController.ChatRequest request = new ChatController.ChatRequest();
request.setId("invalid-chat-session");
request.setQuestion(" ");
ResponseEntity<ChatController.ApiResponse<ChatController.ChatResponse>> response = controller.chat(request);
ChatController.ChatResponse body = response.getBody().getData();
assertFalse(body.isSuccess());
assertEquals("问题内容不能为空", body.getErrorMessage());
verifyNoInteractions(chatService);
}
}
@@ -7,10 +7,12 @@ import com.superbiz.agent.agent.tool.DateTimeTools;
import com.superbiz.agent.agent.tool.QueryLogsTools;
import com.superbiz.agent.agent.tool.QueryMetricsTools;
import com.superbiz.agent.domain.entity.AgentStep;
import com.superbiz.agent.domain.entity.DiagnosisSession;
import com.superbiz.agent.domain.entity.ChatSession;
import com.superbiz.agent.domain.entity.DiagnosisRun;
import com.superbiz.agent.domain.entity.ToolInvocation;
import com.superbiz.agent.repository.AgentStepRepository;
import com.superbiz.agent.repository.DiagnosisSessionRepository;
import com.superbiz.agent.repository.ChatSessionRepository;
import com.superbiz.agent.repository.DiagnosisRunRepository;
import com.superbiz.agent.repository.ToolInvocationRepository;
import com.superbiz.agent.tool.LookupKnowledgeTool;
import com.superbiz.agent.tool.RetrievedDocTracker;
@@ -31,11 +33,15 @@ import java.util.concurrent.atomic.AtomicInteger;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertNotEquals;
import static org.junit.jupiter.api.Assertions.assertSame;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.anyString;
import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.ArgumentMatchers.isNull;
import static org.mockito.Mockito.atLeast;
import static org.mockito.Mockito.atLeastOnce;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.when;
@@ -58,8 +64,73 @@ class ChatServiceSequentialAgentTest {
assertTrue(result.answer().contains("连接池 active 达到上限"));
assertFalse(result.answer().contains("\"answer_version\""));
assertEquals("sequential-test-session", result.sessionId());
assertTrue(result.runId().startsWith("run-"));
assertEquals(List.of("chat_planner", "chat_executor", "chat_verifier", "chat_composer"), chatModel.agentCalls);
assertTrue(chatModel.sawVerifierPrompt);
ChatSessionRepository chatSessionRepository =
(ChatSessionRepository) ReflectionTestUtils.getField(chatService, "chatSessionRepository");
DiagnosisRunRepository diagnosisRunRepository =
(DiagnosisRunRepository) ReflectionTestUtils.getField(chatService, "diagnosisRunRepository");
EvaluationService evaluationService =
(EvaluationService) ReflectionTestUtils.getField(chatService, "evaluationService");
ArgumentCaptor<ChatSession> chatSessionCaptor = ArgumentCaptor.forClass(ChatSession.class);
verify(chatSessionRepository, atLeastOnce()).save(chatSessionCaptor.capture());
assertEquals("sequential-test-session", chatSessionCaptor.getValue().getSessionId());
ArgumentCaptor<DiagnosisRun> runCaptor = ArgumentCaptor.forClass(DiagnosisRun.class);
verify(diagnosisRunRepository, atLeastOnce()).save(runCaptor.capture());
DiagnosisRun savedRun = runCaptor.getValue();
assertEquals(result.runId(), savedRun.getRunId());
assertEquals("sequential-test-session", savedRun.getSessionId());
assertEquals("SUCCESS", savedRun.getStatus());
assertEquals(result.answer(), savedRun.getAnswer());
verify(evaluationService).evaluateRun(eq(result.runId()), eq(result.answer()));
}
@Test
void executeChatComplexCreatesDistinctRunsForSameSessionAcrossTurns() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel firstRoundModel = new ScriptedChatModel();
ScriptedChatModel secondRoundModel = new ScriptedChatModel();
String sessionId = "sequential-same-session";
ChatService.ChatResult first = chatService.executeChatComplex(
firstRoundModel,
new ToolCallback[0],
"第一轮:请分析支付超时",
List.of(),
sessionId
);
ChatService.ChatResult second = chatService.executeChatComplex(
secondRoundModel,
new ToolCallback[0],
"第二轮:基于上一轮结论列出缺失证据",
List.of(
Map.of("role", "user", "content", "第一轮:请分析支付超时"),
Map.of("role", "assistant", "content", first.answer())
),
sessionId
);
assertEquals(sessionId, first.sessionId());
assertEquals(sessionId, second.sessionId());
assertNotEquals(first.runId(), second.runId());
DiagnosisRunRepository diagnosisRunRepository =
(DiagnosisRunRepository) ReflectionTestUtils.getField(chatService, "diagnosisRunRepository");
ArgumentCaptor<DiagnosisRun> runCaptor = ArgumentCaptor.forClass(DiagnosisRun.class);
verify(diagnosisRunRepository, atLeast(2)).save(runCaptor.capture());
List<String> savedRunIds = runCaptor.getAllValues().stream()
.filter(run -> sessionId.equals(run.getSessionId()))
.map(DiagnosisRun::getRunId)
.distinct()
.toList();
assertEquals(2, savedRunIds.size());
assertTrue(savedRunIds.contains(first.runId()));
assertTrue(savedRunIds.contains(second.runId()));
}
@Test
@@ -645,9 +716,12 @@ class ChatServiceSequentialAgentTest {
private ChatService createChatService() {
ChatService chatService = new ChatService();
DiagnosisSessionRepository diagnosisSessionRepository = mock(DiagnosisSessionRepository.class);
when(diagnosisSessionRepository.findBySessionId(anyString())).thenReturn(Optional.empty());
when(diagnosisSessionRepository.save(any(DiagnosisSession.class))).thenAnswer(invocation -> invocation.getArgument(0));
ChatSessionRepository chatSessionRepository = mock(ChatSessionRepository.class);
when(chatSessionRepository.findBySessionId(anyString())).thenReturn(Optional.empty());
when(chatSessionRepository.save(any(ChatSession.class))).thenAnswer(invocation -> invocation.getArgument(0));
DiagnosisRunRepository diagnosisRunRepository = mock(DiagnosisRunRepository.class);
when(diagnosisRunRepository.save(any(DiagnosisRun.class))).thenAnswer(invocation -> invocation.getArgument(0));
AtomicInteger stepId = new AtomicInteger(1);
AgentStepRepository agentStepRepository = mock(AgentStepRepository.class);
@@ -660,13 +734,20 @@ class ChatServiceSequentialAgentTest {
});
when(agentStepRepository.findById(any())).thenReturn(Optional.of(new AgentStep()));
when(agentStepRepository.findBySessionIdOrderByStepIndex(anyString())).thenReturn(List.of());
when(agentStepRepository.findByRunIdOrderByStepIndex(anyString())).thenReturn(List.of());
ToolInvocationRepository toolInvocationRepository = mock(ToolInvocationRepository.class);
when(toolInvocationRepository.countBySessionId(anyString())).thenReturn(0L);
when(toolInvocationRepository.countByRunId(anyString())).thenReturn(0L);
when(toolInvocationRepository.findBySessionIdOrderByIdAsc(anyString())).thenReturn(List.of(ToolInvocation.builder()
.id(101L)
.toolName("query_metrics")
.retrievalDetails(evidenceRefs("$.alerts[0]", "active=50 max=50"))
.build()));
when(toolInvocationRepository.findByRunIdOrderByIdAsc(anyString())).thenReturn(List.of(ToolInvocation.builder()
.id(101L)
.toolName("query_metrics")
.retrievalDetails(evidenceRefs("$.alerts[0]", "active=50 max=50"))
.build()));
EvaluationService evaluationService = mock(EvaluationService.class);
RetrievedDocTracker retrievedDocTracker = mock(RetrievedDocTracker.class);
@@ -674,6 +755,7 @@ class ChatServiceSequentialAgentTest {
when(knowledgeDomainService.buildKnowledgeMap()).thenReturn("");
ToolTraceSummaryService toolTraceSummaryService = mock(ToolTraceSummaryService.class);
when(toolTraceSummaryService.buildVerifierTraceSummary(anyString(), anyString())).thenReturn(List.of());
when(toolTraceSummaryService.buildVerifierTraceSummaryForRun(anyString(), anyString())).thenReturn(List.of());
SelfEvaluationMergeService selfEvaluationMergeService = mock(SelfEvaluationMergeService.class);
when(selfEvaluationMergeService.mergeVerifierEvaluation(any(), any())).thenReturn("{}");
ExecutorGatekeeperService executorGatekeeperService = new ExecutorGatekeeperService(toolInvocationRepository);
@@ -681,7 +763,8 @@ class ChatServiceSequentialAgentTest {
ReflectionTestUtils.setField(chatService, "dateTimeTools", new DateTimeTools());
ReflectionTestUtils.setField(chatService, "lookupKnowledgeTool", new LookupKnowledgeTool());
ReflectionTestUtils.setField(chatService, "queryLogsTools", new QueryLogsTools(mock(ToolInvocationRecorder.class)));
ReflectionTestUtils.setField(chatService, "diagnosisSessionRepository", diagnosisSessionRepository);
ReflectionTestUtils.setField(chatService, "chatSessionRepository", chatSessionRepository);
ReflectionTestUtils.setField(chatService, "diagnosisRunRepository", diagnosisRunRepository);
ReflectionTestUtils.setField(chatService, "agentStepRepository", agentStepRepository);
ReflectionTestUtils.setField(chatService, "toolInvocationRepository", toolInvocationRepository);
ReflectionTestUtils.setField(chatService, "evaluationService", evaluationService);
@@ -11,10 +11,29 @@ import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
class ExecutorGatekeeperServiceTest {
@Test
void validateRunUsesRunScopedToolRows() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
when(repository.findByRunIdOrderByIdAsc("run-gatekeeper-1")).thenReturn(List.of(
invocation(101L, "query_metrics", "$.alerts[0]",
"HighCPUUsage firing, service=payment-service, current=92%")
));
ExecutorGatekeeperService service = new ExecutorGatekeeperService(repository);
Map<String, Object> result = service.validateRun("run-gatekeeper-1",
validOutput(101L, "query_metrics", "$.alerts[0]",
"HighCPUUsage firing, service=payment-service, current=92%"),
Map.of("status", "valid"));
assertEquals("pass", result.get("status"));
verify(repository).findByRunIdOrderByIdAsc("run-gatekeeper-1");
}
@Test
void ruleCatalogLoadsDefaultMetadata() {
GatekeeperRuleCatalog catalog = GatekeeperRuleCatalog.loadDefault(new com.fasterxml.jackson.databind.ObjectMapper());
@@ -26,6 +26,36 @@ class ToolInvocationRecorderTest {
private final ObjectMapper objectMapper = new ObjectMapper();
@Test
void recordEvidenceToolWritesRunIdFromExecutionContext() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
when(repository.save(any(ToolInvocation.class))).thenAnswer(invocation -> invocation.getArgument(0));
ToolInvocationRecorder recorder = new ToolInvocationRecorder(repository, new ObjectMapper());
SessionContextHolder.setContext("recorder-run-session", "run-recorder-1");
try {
recorder.recordEvidenceTool(
"query_metrics",
Map.of("query", "active_prometheus_alerts"),
"{\"success\":true,\"alerts\":[]}",
true,
12,
null,
"prometheus_alerts",
ToolInvocationRecorder.EVIDENCE_STATUS_NO_EVIDENCE,
Map.of("metric_family", "prometheus_alerts")
);
} finally {
SessionContextHolder.clear();
}
ArgumentCaptor<ToolInvocation> captor = ArgumentCaptor.forClass(ToolInvocation.class);
verify(repository).save(captor.capture());
ToolInvocation saved = captor.getValue();
assertEquals("recorder-run-session", saved.getSessionId());
assertEquals("run-recorder-1", saved.getRunId());
}
@Test
void recordEvidenceToolPreservesNoEvidenceSemantics() throws Exception {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
@@ -15,6 +15,32 @@ import static org.mockito.Mockito.when;
class ToolTraceSummaryServiceTest {
@Test
void buildVerifierTraceSummaryForRunUsesRunScopedToolRows() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
when(repository.findByRunIdOrderByIdAsc("run-summary-1")).thenReturn(List.of(
ToolInvocation.builder()
.id(101L)
.sessionId("session-1")
.runId("run-summary-1")
.toolName("query_metrics")
.inputParams("{\"query\":\"active_prometheus_alerts\"}")
.outputPreview("active=50 max=50")
.retrievalDetails("{\"retrieved_domains\":[\"prometheus_alerts\"],\"evidence_status\":\"supported\"}")
.success(true)
.build()
));
ToolTraceSummaryService service = new ToolTraceSummaryService(repository);
List<Map<String, Object>> summaries = service.buildVerifierTraceSummaryForRun(
"run-summary-1", "active=50 max=50");
assertEquals(1, summaries.size());
assertEquals("query_metrics", summaries.get(0).get("tool_name"));
assertEquals(List.of(101L), summaries.get(0).get("source_invocation_ids"));
}
@Test
void buildVerifierTraceSummaryTreatsNoEvidenceAsGapWithoutLosingSuccessfulEvidence() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);