feat(harness): improve trace fallback and reasoning audit

This commit is contained in:
zhuyongxin
2026-07-23 17:52:01 +08:00
parent 8fbc443f76
commit e20249c5d9
38 changed files with 1355 additions and 75 deletions
@@ -3,8 +3,10 @@ package com.superbiz.agent.harness.audit;
import com.alibaba.cloud.ai.graph.RunnableConfig;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.domain.entity.AgentStep;
import com.superbiz.agent.domain.entity.AgentReasoningAudit;
import com.superbiz.agent.harness.agent.DiagnosisAgentFactory;
import com.superbiz.agent.repository.AgentStepRepository;
import com.superbiz.agent.repository.AgentReasoningAuditRepository;
import org.junit.jupiter.api.Test;
import org.mockito.ArgumentCaptor;
import org.springframework.ai.chat.messages.AssistantMessage;
@@ -12,10 +14,13 @@ import org.springframework.ai.chat.messages.UserMessage;
import java.util.List;
import java.util.Optional;
import java.util.concurrent.atomic.AtomicReference;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertNull;
import static org.junit.jupiter.api.Assertions.assertNotNull;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.never;
@@ -25,13 +30,16 @@ import static org.mockito.Mockito.when;
class HarnessAgentAuditHookTest {
@Test
void persistsExactIdentityAndMetadataWithoutContentOrArguments() {
void persistsMetadataAndReasoningInSeparateAudit() {
AgentStepRepository repository = mock(AgentStepRepository.class);
AgentReasoningAuditRepository reasoningRepository = mock(AgentReasoningAuditRepository.class);
AgentStep persisted = AgentStep.builder().id(7L).build();
when(repository.save(any(AgentStep.class))).thenReturn(persisted);
when(repository.findById(7L)).thenReturn(Optional.of(persisted));
AtomicReference<DiagnosisTraceAuditEvent> trace = new AtomicReference<>();
HarnessAgentAuditHook hook = new HarnessAgentAuditHook(
repository, new ObjectMapper(), DiagnosisAgentFactory.AGENT_NAME);
repository, new ObjectMapper(), DiagnosisAgentFactory.AGENT_NAME,
trace::set, reasoningRepository);
RunnableConfig config = RunnableConfig.builder()
.addMetadata("sessionId", "session-audit")
.addMetadata("runId", "run-audit")
@@ -40,6 +48,8 @@ class HarnessAgentAuditHookTest {
hook.beforeModel(List.of(new UserMessage("secret-query")), config);
AssistantMessage response = AssistantMessage.builder()
.content("secret-model-output")
.properties(java.util.Map.of(
"reasoning_content", "inspect bounded evidence before selecting query_logs"))
.toolCalls(List.of(new AssistantMessage.ToolCall(
"call-1", "function", "query_logs", "{\"query\":\"secret-argument\"}")))
.build();
@@ -54,8 +64,28 @@ class HarnessAgentAuditHookTest {
assertFalse(started.getModelInput().contains("secret-query"));
assertFalse(completed.getModelOutput().contains("secret-model-output"));
assertFalse(completed.getModelOutput().contains("secret-argument"));
assertEquals("{\"has_text\":true,\"tool_names\":[\"query_logs\"]}", completed.getModelOutput());
assertEquals("{\"has_text\":true,\"tool_names\":[\"query_logs\"],"
+ "\"reasoning_available\":true,\"reasoning_bytes\":52}",
completed.getModelOutput());
assertNull(completed.getThought());
ArgumentCaptor<AgentReasoningAudit> reasoningCaptor =
ArgumentCaptor.forClass(AgentReasoningAudit.class);
verify(reasoningRepository).save(reasoningCaptor.capture());
AgentReasoningAudit reasoning = reasoningCaptor.getValue();
assertEquals("session-audit", reasoning.getSessionId());
assertEquals("run-audit", reasoning.getRunId());
assertEquals(0, reasoning.getStepIndex());
assertTrue(reasoning.getReasoningAvailable());
assertEquals("inspect bounded evidence before selecting query_logs",
reasoning.getReasoningContent());
assertNotNull(trace.get());
assertEquals(TraceEventType.AGENT_MODEL_STEP, trace.get().eventType());
String traceText = trace.get().details().toString();
assertFalse(traceText.contains("secret-query"));
assertFalse(traceText.contains("secret-model-output"));
assertFalse(traceText.contains("secret-argument"));
assertFalse(traceText.contains("inspect bounded evidence"));
assertTrue(traceText.contains("reasoning_available=true"));
}
@Test
@@ -0,0 +1,78 @@
package com.superbiz.agent.harness.audit;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.domain.entity.DiagnosisTraceEvent;
import com.superbiz.agent.repository.DiagnosisTraceEventRepository;
import org.junit.jupiter.api.Test;
import org.mockito.ArgumentCaptor;
import org.springframework.core.io.ClassPathResource;
import java.nio.charset.StandardCharsets;
import java.util.Map;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertDoesNotThrow;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.Mockito.doThrow;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.times;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
class JpaDiagnosisTraceRecorderTest {
@Test
void persistsIndependentMonotonicSafeEvents() {
DiagnosisTraceEventRepository repository = mock(DiagnosisTraceEventRepository.class);
when(repository.findMaxSequenceNoByRunId("run-1")).thenReturn(7);
JpaDiagnosisTraceRecorder recorder = new JpaDiagnosisTraceRecorder(
repository, new ObjectMapper());
recorder.record(event(TraceEventType.RUN_STARTED, TraceEventStatus.STARTED));
recorder.record(event(TraceEventType.RELEASE_DECISION, TraceEventStatus.FALLBACK));
ArgumentCaptor<DiagnosisTraceEvent> captor =
ArgumentCaptor.forClass(DiagnosisTraceEvent.class);
verify(repository, times(2)).save(captor.capture());
assertEquals(8, captor.getAllValues().get(0).getSequenceNo());
assertEquals(9, captor.getAllValues().get(1).getSequenceNo());
assertEquals("{\"failure\":\"SCHEMA_INVALID\"}",
captor.getAllValues().get(0).getDetails());
assertFalse(captor.getAllValues().get(0).getDetails().contains("prompt"));
}
@Test
void persistenceFailureNeverChangesMainFlow() {
DiagnosisTraceEventRepository repository = mock(DiagnosisTraceEventRepository.class);
when(repository.findMaxSequenceNoByRunId("run-1")).thenReturn(0);
doThrow(new IllegalStateException("database unavailable"))
.when(repository).save(any(DiagnosisTraceEvent.class));
JpaDiagnosisTraceRecorder recorder = new JpaDiagnosisTraceRecorder(
repository, new ObjectMapper());
assertDoesNotThrow(() -> recorder.record(
event(TraceEventType.RUN_STARTED, TraceEventStatus.STARTED)));
}
@Test
void migrationCreatesIndependentAppendOnlyTraceTable() throws Exception {
String sql;
try (var input = new ClassPathResource(
"db/migration/V014__create_diagnosis_trace_event.sql").getInputStream()) {
sql = new String(input.readAllBytes(), StandardCharsets.UTF_8);
}
assertTrue(sql.contains("CREATE TABLE diagnosis_trace_event"));
assertTrue(sql.contains("details JSON NOT NULL"));
assertTrue(sql.contains("idx_trace_event_run_sequence"));
assertFalse(sql.toUpperCase().contains("FOREIGN KEY"));
}
private DiagnosisTraceAuditEvent event(TraceEventType type, TraceEventStatus status) {
return new DiagnosisTraceAuditEvent(
"session-1", "run-1", TracePhase.RUN, type, status,
null, null, Map.of("failure", "SCHEMA_INVALID"));
}
}
@@ -18,7 +18,9 @@ class JpaToolInvocationAuditSinkTest {
@Test
void persistsOnlyBoundedStableMetadata() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
JpaToolInvocationAuditSink sink = new JpaToolInvocationAuditSink(repository, new ObjectMapper());
DiagnosisTraceRecorder traceRecorder = mock(DiagnosisTraceRecorder.class);
JpaToolInvocationAuditSink sink = new JpaToolInvocationAuditSink(
repository, new ObjectMapper(), traceRecorder);
sink.record(new ToolInvocationAuditEvent(
"session-1", "run-1", "call-1", "query_logs",
@@ -37,5 +39,11 @@ class JpaToolInvocationAuditSinkTest {
String serialized = saved.getInputParams() + saved.getOutputPreview() + saved.getRetrievalDetails();
assertFalse(serialized.contains("query"));
assertFalse(serialized.contains("raw_response"));
ArgumentCaptor<DiagnosisTraceAuditEvent> traceCaptor =
ArgumentCaptor.forClass(DiagnosisTraceAuditEvent.class);
verify(traceRecorder).record(traceCaptor.capture());
assertEquals(TraceEventType.TOOL_INVOCATION, traceCaptor.getValue().eventType());
assertEquals("call-1", traceCaptor.getValue().details().get("tool_call_id"));
assertFalse(traceCaptor.getValue().details().toString().contains("raw_response"));
}
}