feat(harness): improve trace fallback and reasoning audit
This commit is contained in:
@@ -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"));
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user