Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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);
}
}
Expand All @@ -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()) {
Expand All @@ -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<String> 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());
}
}
Original file line number Diff line number Diff line change
@@ -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<String> responses;
private final List<String> 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;
}
}
}
Loading