diff --git a/src/main/java/com/springclaw/service/memory/extraction/ProviderBackedTerminalReflectionGenerator.java b/src/main/java/com/springclaw/service/memory/extraction/ProviderBackedTerminalReflectionGenerator.java index ebc3203f..02a729c1 100644 --- a/src/main/java/com/springclaw/service/memory/extraction/ProviderBackedTerminalReflectionGenerator.java +++ b/src/main/java/com/springclaw/service/memory/extraction/ProviderBackedTerminalReflectionGenerator.java @@ -29,9 +29,15 @@ class ProviderBackedTerminalReflectionGenerator implements TerminalReflectionGen public TerminalReflectionResult reflect(TerminalMemoryExtractionContext context) throws Exception { String first = call(context, renderPrompt(context)); try { - return jsonParser.parseObject(first, TerminalReflectionResult.class); + TerminalReflectionResult result = + jsonParser.parseObject(first, TerminalReflectionResult.class); + if (validReflection(context, result)) { + return result; + } + String repaired = call(context, renderSemanticRepairPrompt(context)); + return jsonParser.parseObject(repaired, TerminalReflectionResult.class); } catch (Exception parseFailure) { - String repaired = call(context, "上一次输出不是合法 JSON。请只输出 terminal-reflection.v1 JSON。"); + String repaired = call(context, renderJsonRepairPrompt(context)); return jsonParser.parseObject(repaired, TerminalReflectionResult.class); } } @@ -58,6 +64,33 @@ private String renderPrompt(TerminalMemoryExtractionContext context) { """.formatted(context.runId(), renderEvidence(context)); } + private String renderJsonRepairPrompt(TerminalMemoryExtractionContext context) { + return """ + 上一次输出不是合法 JSON。请只输出 springclaw.terminal-reflection.v1 JSON 对象。 + lesson 不能为空;evidenceRefs 必须包含 run:%s,并至少包含一个下列 event 引用: + %s + 不要 Markdown,不要解释。 + """.formatted(context.runId(), renderAllowedEvidence(context)); + } + + private String renderSemanticRepairPrompt(TerminalMemoryExtractionContext context) { + return """ + 上一次输出 JSON 合法,但反思内容未通过校验。 + 请只输出 springclaw.terminal-reflection.v1 JSON 对象,并满足: + - lesson 不能为空,必须是可复用教训; + - evidenceRefs 必须包含 run:%s; + - evidenceRefs 至少包含一个下列 event 引用; + - 不允许使用未列出的证据引用; + - 没有可复用教训时,也要用一句保守 lesson 说明“本轮没有可推广教训”并引用证据。 + + 可用证据引用: + %s + + 原始证据: + %s + """.formatted(context.runId(), renderAllowedEvidence(context), renderEvidence(context)); + } + private String renderEvidence(TerminalMemoryExtractionContext context) { StringBuilder sb = new StringBuilder(); for (MemorySourceEvent event : context.events()) { @@ -68,4 +101,38 @@ private String renderEvidence(TerminalMemoryExtractionContext context) { } return sb.toString(); } + + private String renderAllowedEvidence(TerminalMemoryExtractionContext context) { + StringBuilder sb = new StringBuilder(); + sb.append("- run:").append(context.runId()).append('\n'); + for (MemorySourceEvent event : context.events()) { + sb.append("- event:").append(event.eventKey()).append('\n'); + } + return sb.toString(); + } + + private boolean validReflection( + TerminalMemoryExtractionContext context, + TerminalReflectionResult reflection + ) { + if (reflection == null + || !TerminalMemoryExtractionService.TERMINAL_REFLECTION_SCHEMA.equals(reflection.schema()) + || reflection.lesson() == null + || reflection.lesson().isBlank() + || !Double.isFinite(reflection.confidence()) + || reflection.confidence() < 0.0 + || reflection.confidence() > 1.0 + || reflection.evidenceRefs() == null + || reflection.evidenceRefs().isEmpty()) { + return false; + } + java.util.Set allowed = new java.util.LinkedHashSet<>(); + allowed.add("run:" + context.runId()); + for (MemorySourceEvent event : context.events()) { + allowed.add("event:" + event.eventKey()); + } + return reflection.evidenceRefs().contains("run:" + context.runId()) + && reflection.evidenceRefs().stream().anyMatch(ref -> ref != null && ref.startsWith("event:")) + && allowed.containsAll(reflection.evidenceRefs()); + } } diff --git a/src/test/java/com/springclaw/service/memory/extraction/ProviderBackedTerminalReflectionGeneratorTest.java b/src/test/java/com/springclaw/service/memory/extraction/ProviderBackedTerminalReflectionGeneratorTest.java new file mode 100644 index 00000000..2af178b8 --- /dev/null +++ b/src/test/java/com/springclaw/service/memory/extraction/ProviderBackedTerminalReflectionGeneratorTest.java @@ -0,0 +1,119 @@ +package com.springclaw.service.memory.extraction; + +import com.fasterxml.jackson.databind.ObjectMapper; +import com.springclaw.service.ai.AiProviderService; +import org.junit.jupiter.api.Test; + +import java.util.ArrayList; +import java.util.List; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.mock; + +class ProviderBackedTerminalReflectionGeneratorTest { + + @Test + void repairsParsedButUngroundedReflectionWithEvidenceSpecificPrompt() throws Exception { + CapturingModelClient modelClient = new CapturingModelClient( + """ + {"schema":"springclaw.terminal-reflection.v1","outcome":"SUCCESS","lesson":"","applicability":"","failureMode":"","evidenceRefs":[],"confidence":0.8} + """, + """ + {"schema":"springclaw.terminal-reflection.v1","outcome":"SUCCESS","lesson":"用户明确要求进度汇报用简短中文时,后续汇报应保持简短中文。","applicability":"后续进度汇报","failureMode":"","evidenceRefs":["run:run-1","event:chat:run-1:user"],"confidence":0.86} + """ + ); + ProviderBackedTerminalReflectionGenerator generator = + new ProviderBackedTerminalReflectionGenerator( + modelClient, + new StrictJsonParser(new ObjectMapper()), + "deepseek", + "deepseek" + ); + + TerminalReflectionResult result = generator.reflect(context()); + + assertThat(result.lesson()).contains("简短中文"); + assertThat(result.evidenceRefs()) + .containsExactly("run:run-1", "event:chat:run-1:user"); + assertThat(modelClient.userPrompts).hasSize(2); + assertThat(modelClient.userPrompts.get(1)) + .contains("上一次输出 JSON 合法,但反思内容未通过校验") + .contains("run:run-1") + .contains("event:chat:run-1:user") + .contains("lesson 不能为空"); + } + + @Test + void repairsReflectionThatOmitsEventGrounding() throws Exception { + CapturingModelClient modelClient = new CapturingModelClient( + """ + {"schema":"springclaw.terminal-reflection.v1","outcome":"SUCCESS","lesson":"后续进度汇报应保持简短中文。","applicability":"后续进度汇报","failureMode":"","evidenceRefs":["run:run-1"],"confidence":0.86} + """, + """ + {"schema":"springclaw.terminal-reflection.v1","outcome":"SUCCESS","lesson":"用户明确要求进度汇报用简短中文时,后续汇报应保持简短中文。","applicability":"后续进度汇报","failureMode":"","evidenceRefs":["run:run-1","event:chat:run-1:user"],"confidence":0.86} + """ + ); + ProviderBackedTerminalReflectionGenerator generator = + new ProviderBackedTerminalReflectionGenerator( + modelClient, + new StrictJsonParser(new ObjectMapper()), + "deepseek", + "deepseek" + ); + + TerminalReflectionResult result = generator.reflect(context()); + + assertThat(result.evidenceRefs()) + .containsExactly("run:run-1", "event:chat:run-1:user"); + assertThat(modelClient.userPrompts).hasSize(2); + } + + private static TerminalMemoryExtractionContext context() { + return new TerminalMemoryExtractionContext( + "run-1", + "session-1", + "api", + "alice", + List.of( + new MemorySourceEvent( + "chat:run-1:user", + "USER", + "CHAT", + "以后给我进度汇报请用简短中文。" + ), + new MemorySourceEvent( + "chat:run-1:assistant:terminal", + "ASSISTANT", + "CHAT", + "收到,后续进度会用简短中文说明。" + ) + ) + ); + } + + private static final class CapturingModelClient extends ProviderMemoryModelClient { + private final List responses; + private final List userPrompts = new ArrayList<>(); + private int index; + + private CapturingModelClient(String... responses) { + super(mock(AiProviderService.class)); + this.responses = List.of(responses); + } + + @Override + String call( + String providerId, + String fallbackProviderId, + String source, + TerminalMemoryExtractionContext context, + String systemPrompt, + String userPrompt + ) { + userPrompts.add(userPrompt); + String response = responses.get(Math.min(index, responses.size() - 1)); + index++; + return response; + } + } +}