diff --git a/.gitignore b/.gitignore index 3dd6d29..9c8bd2c 100644 --- a/.gitignore +++ b/.gitignore @@ -11,6 +11,8 @@ build # Environment variables .env test-output.txt +# 策略评测报告输出 +eval-output/ # IDEA code analysis qodana.yaml diff --git a/app/src/main/java/com/quantmore/modules/generator/dto/GenerateStrategyRequest.java b/app/src/main/java/com/quantmore/modules/generator/dto/GenerateStrategyRequest.java index 50ab7e8..ad42a6f 100644 --- a/app/src/main/java/com/quantmore/modules/generator/dto/GenerateStrategyRequest.java +++ b/app/src/main/java/com/quantmore/modules/generator/dto/GenerateStrategyRequest.java @@ -32,6 +32,9 @@ public record GenerateStrategyRequest( List knowledgeBaseIds, // 可选:空 = 用户默认模型 - String providerId + String providerId, + + // 可选:true = 跳过知识库检索(无 RAG 对照生成),null/false = 正常检索 + Boolean skipRetrieval ) { } diff --git a/app/src/main/java/com/quantmore/modules/generator/eval/EvalCase.java b/app/src/main/java/com/quantmore/modules/generator/eval/EvalCase.java new file mode 100644 index 0000000..e882296 --- /dev/null +++ b/app/src/main/java/com/quantmore/modules/generator/eval/EvalCase.java @@ -0,0 +1,16 @@ +package com.quantmore.modules.generator.eval; + +/** + * 评测用例(一条策略需求描述) + */ +public record EvalCase( + String id, + String name, + String market, + String frequency, + String buyConditions, + String sellConditions, + String riskControls, + String difficulty +) { +} diff --git a/app/src/main/java/com/quantmore/modules/generator/eval/EvalCaseLoader.java b/app/src/main/java/com/quantmore/modules/generator/eval/EvalCaseLoader.java new file mode 100644 index 0000000..77265d5 --- /dev/null +++ b/app/src/main/java/com/quantmore/modules/generator/eval/EvalCaseLoader.java @@ -0,0 +1,71 @@ +package com.quantmore.modules.generator.eval; + +import com.fasterxml.jackson.core.type.TypeReference; +import com.fasterxml.jackson.databind.ObjectMapper; +import lombok.RequiredArgsConstructor; +import org.springframework.core.io.Resource; +import org.springframework.core.io.ResourceLoader; +import org.springframework.stereotype.Component; + +import java.io.IOException; +import java.io.InputStream; +import java.util.List; +import java.util.Set; + +/** + * 评测用例加载与校验:从 casesPath 读取 JSON 数组并校验字段合法性 + */ +@Component +@RequiredArgsConstructor +public class EvalCaseLoader { + + private static final Set MARKETS = Set.of("STOCK", "ETF", "CONVERTIBLE_BOND", "FUTURES"); + private static final Set FREQUENCIES = Set.of("DAILY", "MINUTE"); + private static final Set DIFFICULTIES = Set.of("SIMPLE", "MEDIUM", "COMPLEX"); + + private final ResourceLoader resourceLoader; + private final ObjectMapper objectMapper; + private final EvalProperties properties; + + public List load() { + Resource resource = resourceLoader.getResource(properties.getCasesPath()); + if (!resource.exists()) { + throw new IllegalStateException("评测用例文件不存在: " + properties.getCasesPath()); + } + List cases; + try (InputStream in = resource.getInputStream()) { + cases = objectMapper.readValue(in, new TypeReference>() { + }); + } catch (IOException e) { + throw new IllegalStateException("评测用例文件解析失败: " + properties.getCasesPath(), e); + } + validate(cases); + return cases; + } + + private void validate(List cases) { + if (cases == null || cases.isEmpty()) { + throw new IllegalStateException("评测用例为空"); + } + for (EvalCase c : cases) { + requireNotBlank(c.id(), "id"); + requireNotBlank(c.name(), "name"); + requireNotBlank(c.buyConditions(), "buyConditions"); + if (!MARKETS.contains(c.market())) { + throw new IllegalStateException("用例 " + c.id() + " market 非法: " + c.market()); + } + if (!FREQUENCIES.contains(c.frequency())) { + throw new IllegalStateException("用例 " + c.id() + " frequency 非法: " + c.frequency()); + } + if (!DIFFICULTIES.contains(c.difficulty())) { + throw new IllegalStateException("用例 " + c.id() + " difficulty 非法: " + c.difficulty()); + } + } + } + + private void requireNotBlank(String value, String field) { + if (value == null || value.isBlank()) { + throw new IllegalStateException("用例字段为空: " + field); + } + } +} diff --git a/app/src/main/java/com/quantmore/modules/generator/eval/EvalJudgeService.java b/app/src/main/java/com/quantmore/modules/generator/eval/EvalJudgeService.java new file mode 100644 index 0000000..03909dc --- /dev/null +++ b/app/src/main/java/com/quantmore/modules/generator/eval/EvalJudgeService.java @@ -0,0 +1,94 @@ +package com.quantmore.modules.generator.eval; + +import com.quantmore.common.ai.LlmProviderRegistry; +import com.quantmore.common.ai.PromptSanitizer; +import com.quantmore.common.ai.PromptSecurityConstants; +import com.quantmore.common.exception.BusinessException; +import com.quantmore.common.exception.ErrorCode; +import lombok.extern.slf4j.Slf4j; +import org.springframework.ai.chat.client.ChatClient; +import org.springframework.ai.chat.prompt.PromptTemplate; +import org.springframework.core.io.ClassPathResource; +import org.springframework.stereotype.Component; + +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.util.List; +import java.util.Map; + +/** + * LLM 评委:按 rubric 对生成代码评分,返回结构化 JSON 结果。 + * 使用 getPlainChatClient(无工具/记忆 advisor),保证输出为可解析的纯文本 JSON。 + */ +@Slf4j +@Component +public class EvalJudgeService { + + private final LlmProviderRegistry registry; + private final PromptSanitizer sanitizer; + private final EvalProperties properties; + private final PromptTemplate template; + + public EvalJudgeService( + LlmProviderRegistry registry, + PromptSanitizer sanitizer, + EvalProperties properties) throws IOException { + this.registry = registry; + this.sanitizer = sanitizer; + this.properties = properties; + this.template = new PromptTemplate( + new ClassPathResource("prompts/strategy-eval-judge.st") + .getContentAsString(StandardCharsets.UTF_8)); + } + + /** + * 评分失败抛 BusinessException(AI_SERVICE_ERROR),由评测服务记 judgeFailed + */ + public JudgeResult judge(EvalCase caseMeta, String code) { + String systemPrompt = template.render(Map.of( + "strategyName", sanitizer.sanitize(caseMeta.name()).trim(), + "market", caseMeta.market(), + "frequency", caseMeta.frequency(), + "buyConditions", sanitizer.sanitize(caseMeta.buyConditions()).trim(), + "sellConditions", sanitizeNullable(caseMeta.sellConditions()), + "riskControls", sanitizeNullable(caseMeta.riskControls()), + "generatedCode", sanitizer.wrapWithDelimiters("generated-code", code) + )) + PromptSecurityConstants.ANTI_INJECTION_INSTRUCTION; + + String raw; + try { + raw = resolveClient().prompt() + .system(systemPrompt) + .call() + .chatClientResponse() + .chatResponse() + .getResult() + .getOutput() + .getText(); + } catch (Exception e) { + log.error("评委评分失败: case={}, error={}", caseMeta.id(), e.getMessage(), e); + throw new BusinessException(ErrorCode.AI_SERVICE_ERROR, "评委评分失败: " + e.getMessage()); + } + if (raw == null || raw.isBlank()) { + throw new BusinessException(ErrorCode.AI_SERVICE_ERROR, "评委返回内容为空"); + } + return JudgeJsonParser.parse(raw); + } + + private String sanitizeNullable(String value) { + return value == null ? "" : sanitizer.sanitize(value).trim(); + } + + private ChatClient resolveClient() { + String providerId = properties.getJudgeProvider(); + return (providerId == null || providerId.isBlank()) + ? registry.getPlainChatClient() + : registry.getPlainChatClient(providerId); + } + + public record JudgeResult(double score, boolean passed, List issues) { + } + + public record JudgeIssue(String dimension, String comment) { + } +} diff --git a/app/src/main/java/com/quantmore/modules/generator/eval/EvalProperties.java b/app/src/main/java/com/quantmore/modules/generator/eval/EvalProperties.java new file mode 100644 index 0000000..cc606d9 --- /dev/null +++ b/app/src/main/java/com/quantmore/modules/generator/eval/EvalProperties.java @@ -0,0 +1,46 @@ +package com.quantmore.modules.generator.eval; + +import lombok.Data; +import org.springframework.boot.context.properties.ConfigurationProperties; +import org.springframework.stereotype.Component; + +import java.time.Duration; + +/** + * 策略生成评测配置(均可用 APP_EVAL_* 环境变量覆盖) + */ +@Data +@Component +@ConfigurationProperties(prefix = "app.eval") +public class EvalProperties { + + /** 是否启用评测:APP_EVAL_ENABLED=true 时 bootRun 跑完评测写报告后自动退出 */ + private boolean enabled = false; + + /** 用例文件路径,支持 classpath: / file: 前缀 */ + private String casesPath = "classpath:eval/strategy-eval-cases.json"; + + /** 生成用 provider id(空 = 用户默认/全局默认) */ + private String generateProvider; + + /** 评委用 provider id(空 = 全局默认) */ + private String judgeProvider; + + /** 报告输出目录(相对 bootRun 工作目录) */ + private String outputDir = "eval-output"; + + /** 评委通过分数线 */ + private double judgePassScore = 70.0; + + /** python 解释器路径 */ + private String pythonBin = "python3"; + + /** 单次 python 语法检查超时 */ + private Duration pythonTimeout = Duration.ofSeconds(10); + + /** 等待知识库向量化就绪的超时 */ + private Duration vectorWaitTimeout = Duration.ofSeconds(120); + + /** 评测结束后是否清理本次评测写入的生成记录 */ + private boolean cleanupRecords = true; +} diff --git a/app/src/main/java/com/quantmore/modules/generator/eval/EvalReport.java b/app/src/main/java/com/quantmore/modules/generator/eval/EvalReport.java new file mode 100644 index 0000000..ac4ddb0 --- /dev/null +++ b/app/src/main/java/com/quantmore/modules/generator/eval/EvalReport.java @@ -0,0 +1,139 @@ +package com.quantmore.modules.generator.eval; + +import java.util.List; + +/** + * 评测报告数据结构(嵌套 record,聚合逻辑见 EvalSummary.of) + */ +public final class EvalReport { + + private EvalReport() { + } + + /** + * 单个分支(RAG 或无 RAG)的执行结果 + */ + public record BranchResult( + boolean ragEnabled, + boolean generationOk, + String generationError, + Long generationId, + long generationMs, + PythonSyntaxCheckService.SyntaxCheckResult syntax, + boolean judgeOk, + EvalJudgeService.JudgeResult judge, + String judgeRaw + ) { + } + + /** + * 单个用例的 RAG / no-RAG 对照结果 + */ + public record CaseResult(EvalCase caseMeta, BranchResult rag, BranchResult noRag) { + } + + /** + * 汇总统计(纯函数聚合,便于单测) + */ + public record EvalSummary( + int totalCases, + int ragPassed, + int noRagPassed, + double ragAvgScore, + double noRagAvgScore, + int ragSyntaxPassed, + int noRagSyntaxPassed, + int generationFailures, + int judgeFailures, + int py35WarningCount, + double scoreDelta + ) { + + public static EvalSummary of(List results, double passScore) { + int ragPassed = 0; + int noRagPassed = 0; + int ragSyntaxPassed = 0; + int noRagSyntaxPassed = 0; + int generationFailures = 0; + int judgeFailures = 0; + int py35WarningCount = 0; + double ragScoreSum = 0; + int ragScoreCount = 0; + double noRagScoreSum = 0; + int noRagScoreCount = 0; + + for (CaseResult result : results) { + ragPassed += passed(result.rag(), passScore) ? 1 : 0; + noRagPassed += passed(result.noRag(), passScore) ? 1 : 0; + ragSyntaxPassed += syntaxPassed(result.rag()) ? 1 : 0; + noRagSyntaxPassed += syntaxPassed(result.noRag()) ? 1 : 0; + generationFailures += result.rag().generationOk() ? 0 : 1; + generationFailures += result.noRag().generationOk() ? 0 : 1; + judgeFailures += judgeFailed(result.rag()) ? 1 : 0; + judgeFailures += judgeFailed(result.noRag()) ? 1 : 0; + py35WarningCount += warnings(result.rag()).size(); + py35WarningCount += warnings(result.noRag()).size(); + if (result.rag().judgeOk()) { + ragScoreSum += result.rag().judge().score(); + ragScoreCount++; + } + if (result.noRag().judgeOk()) { + noRagScoreSum += result.noRag().judge().score(); + noRagScoreCount++; + } + } + + double ragAvgScore = ragScoreCount == 0 ? 0 : ragScoreSum / ragScoreCount; + double noRagAvgScore = noRagScoreCount == 0 ? 0 : noRagScoreSum / noRagScoreCount; + return new EvalSummary( + results.size(), + ragPassed, + noRagPassed, + ragAvgScore, + noRagAvgScore, + ragSyntaxPassed, + noRagSyntaxPassed, + generationFailures, + judgeFailures, + py35WarningCount, + ragAvgScore - noRagAvgScore + ); + } + + private static boolean passed(BranchResult branch, double passScore) { + return branch.generationOk() + && syntaxPassed(branch) + && branch.judgeOk() + && branch.judge().score() >= passScore; + } + + private static boolean syntaxPassed(BranchResult branch) { + return branch.generationOk() && "PASS".equals(branch.syntax().status()); + } + + private static boolean judgeFailed(BranchResult branch) { + return branch.generationOk() && !branch.judgeOk(); + } + + private static List warnings(BranchResult branch) { + return branch.syntax() == null ? List.of() : branch.syntax().py35Warnings(); + } + } + + /** + * 完整评测报告 + */ + public record FullReport( + String runAt, + String evalUser, + String generateProvider, + String judgeProvider, + String pythonVersion, + boolean kbReady, + int kbCompleted, + int kbFailed, + List results, + EvalSummary summary + ) { + } +} diff --git a/app/src/main/java/com/quantmore/modules/generator/eval/EvalReportWriter.java b/app/src/main/java/com/quantmore/modules/generator/eval/EvalReportWriter.java new file mode 100644 index 0000000..d06ef36 --- /dev/null +++ b/app/src/main/java/com/quantmore/modules/generator/eval/EvalReportWriter.java @@ -0,0 +1,148 @@ +package com.quantmore.modules.generator.eval; + +import com.fasterxml.jackson.databind.ObjectMapper; +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; +import org.springframework.stereotype.Component; + +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; + +/** + * 评测报告输出:eval-output/ 下写 eval-report.json 与 eval-report.md,并打控制台汇总。 + */ +@Slf4j +@Component +@RequiredArgsConstructor +public class EvalReportWriter { + + private final EvalProperties properties; + private final ObjectMapper objectMapper; + + public void write(EvalReport.FullReport report) throws IOException { + Path dir = Path.of(properties.getOutputDir()); + Files.createDirectories(dir); + objectMapper.writerWithDefaultPrettyPrinter() + .writeValue(dir.resolve("eval-report.json").toFile(), report); + Files.writeString(dir.resolve("eval-report.md"), toMarkdown(report), StandardCharsets.UTF_8); + log.info("评测报告已写出: dir={}", dir.toAbsolutePath()); + printConsole(report); + } + + private String toMarkdown(EvalReport.FullReport report) { + EvalReport.EvalSummary s = report.summary(); + StringBuilder sb = new StringBuilder(); + sb.append("# 策略生成评测报告\n\n"); + sb.append("- 运行时间: ").append(report.runAt()).append('\n'); + sb.append("- 评测用户: ").append(report.evalUser()).append('\n'); + sb.append("- 生成模型: ").append(display(report.generateProvider())).append('\n'); + sb.append("- 评委模型: ").append(display(report.judgeProvider())).append('\n'); + sb.append("- Python: ").append(display(report.pythonVersion())).append('\n'); + sb.append("- 知识库就绪: ").append(report.kbReady() ? "是" : "否") + .append("(COMPLETED ").append(report.kbCompleted()) + .append(" / FAILED ").append(report.kbFailed()).append(")\n\n"); + + sb.append("## 汇总(通过 = 语法 PASS 且评委分 >= ") + .append(properties.getJudgePassScore()).append(")\n\n"); + sb.append("| 指标 | RAG | no-RAG |\n|---|---|---|\n"); + sb.append("| 用例通过数 | ").append(s.ragPassed()).append(" | ").append(s.noRagPassed()) + .append(" |\n"); + sb.append("| 语法通过数 | ").append(s.ragSyntaxPassed()).append(" | ") + .append(s.noRagSyntaxPassed()).append(" |\n"); + sb.append("| 评委平均分 | ").append(String.format("%.1f", s.ragAvgScore())).append(" | ") + .append(String.format("%.1f", s.noRagAvgScore())).append(" |\n"); + sb.append("\n- 总用例数: ").append(s.totalCases()).append('\n'); + sb.append("- 生成失败分支: ").append(s.generationFailures()).append('\n'); + sb.append("- 评委失败分支: ").append(s.judgeFailures()).append('\n'); + sb.append("- Python 3.5 兼容警示数: ").append(s.py35WarningCount()).append('\n'); + sb.append("- 平均分差(RAG - no-RAG): ").append(String.format("%.1f", s.scoreDelta())) + .append("\n\n"); + + sb.append("## 逐用例\n\n"); + sb.append("| 用例 | 难度 | RAG 语法 | RAG 评分 | RAG 通过 | no-RAG 语法 | no-RAG 评分 | no-RAG 通过 |\n"); + sb.append("|---|---|---|---|---|---|---|---|\n"); + for (EvalReport.CaseResult r : report.results()) { + sb.append("| ").append(r.caseMeta().id()).append(" ").append(r.caseMeta().name()) + .append(" | ").append(r.caseMeta().difficulty()) + .append(" | ").append(branchCell(r.rag())) + .append(" |\n"); + } + + sb.append("\n## 问题明细\n\n"); + boolean hasIssues = false; + for (EvalReport.CaseResult r : report.results()) { + appendBranchIssues(sb, r.caseMeta(), r.rag()); + appendBranchIssues(sb, r.caseMeta(), r.noRag()); + hasIssues = hasIssues || !r.rag().judgeOk() || !r.noRag().judgeOk() + || !"PASS".equals(syntax(r.rag())) || !"PASS".equals(syntax(r.noRag())); + } + if (!hasIssues) { + sb.append("无\n"); + } + return sb.toString(); + } + + private void appendBranchIssues(StringBuilder sb, EvalCase caseMeta, + EvalReport.BranchResult branch) { + String label = branch.ragEnabled() ? "RAG" : "no-RAG"; + if (!branch.generationOk()) { + sb.append("- ").append(caseMeta.id()).append(" ").append(label) + .append(" 生成失败: ").append(branch.generationError()).append('\n'); + return; + } + if (!"PASS".equals(syntax(branch))) { + sb.append("- ").append(caseMeta.id()).append(" ").append(label) + .append(" 语法 ").append(syntax(branch)).append(": ") + .append(branch.syntax() == null ? "" : branch.syntax().message()).append('\n'); + } + if (!branch.judgeOk()) { + sb.append("- ").append(caseMeta.id()).append(" ").append(label) + .append(" 评委失败: ").append(branch.judgeRaw()).append('\n'); + } else if (branch.judge().issues() != null && !branch.judge().issues().isEmpty()) { + for (EvalJudgeService.JudgeIssue issue : branch.judge().issues()) { + sb.append("- ").append(caseMeta.id()).append(" ").append(label) + .append(" [").append(issue.dimension()).append("] ") + .append(issue.comment()).append('\n'); + } + } + List warnings = branch.syntax() == null ? List.of() : branch.syntax().py35Warnings(); + for (String warning : warnings) { + sb.append("- ").append(caseMeta.id()).append(" ").append(label) + .append(" 3.5 兼容警示: ").append(warning).append('\n'); + } + } + + private String branchCell(EvalReport.BranchResult branch) { + String syntax = branch.generationOk() ? syntax(branch) : "生成失败"; + String score = !branch.generationOk() ? "-" + : branch.judgeOk() ? String.format("%.0f", branch.judge().score()) : "评委失败"; + boolean passed = branch.generationOk() && "PASS".equals(syntax(branch)) + && branch.judgeOk() && branch.judge().score() >= properties.getJudgePassScore(); + return syntax + " | " + score + " | " + (passed ? "是" : "否"); + } + + private String syntax(EvalReport.BranchResult branch) { + return branch.syntax() == null ? "-" : branch.syntax().status(); + } + + private String display(String value) { + return (value == null || value.isBlank()) ? "(默认)" : value; + } + + private void printConsole(EvalReport.FullReport report) { + EvalReport.EvalSummary s = report.summary(); + log.info("================ 评测汇总 ================"); + log.info("用例数={} 通过: RAG={}/{} no-RAG={}/{}", s.totalCases(), s.ragPassed(), + s.totalCases(), s.noRagPassed(), s.totalCases()); + log.info("语法通过: RAG={}/{} no-RAG={}/{}", s.ragSyntaxPassed(), s.totalCases(), + s.noRagSyntaxPassed(), s.totalCases()); + log.info("平均分: RAG={} no-RAG={} 差值={}", String.format("%.1f", s.ragAvgScore()), + String.format("%.1f", s.noRagAvgScore()), String.format("%+.1f", s.scoreDelta())); + log.info("生成失败={} 评委失败={} 3.5警示={}", s.generationFailures(), s.judgeFailures(), + s.py35WarningCount()); + log.info("========================================="); + } +} diff --git a/app/src/main/java/com/quantmore/modules/generator/eval/JudgeJsonParser.java b/app/src/main/java/com/quantmore/modules/generator/eval/JudgeJsonParser.java new file mode 100644 index 0000000..9e3194f --- /dev/null +++ b/app/src/main/java/com/quantmore/modules/generator/eval/JudgeJsonParser.java @@ -0,0 +1,66 @@ +package com.quantmore.modules.generator.eval; + +import com.fasterxml.jackson.databind.DeserializationFeature; +import com.fasterxml.jackson.databind.ObjectMapper; + +import java.io.IOException; +import java.util.List; +import java.util.regex.Matcher; +import java.util.regex.Pattern; + +/** + * 评委输出解析:优先提取 ```json 围栏内容,否则取首尾大括号子串,Jackson 反序列化。 + * 解析失败抛 IllegalStateException(原文由调用方留档)。 + */ +public final class JudgeJsonParser { + + private static final Pattern JSON_FENCE = Pattern.compile("(?s)```(?:json)?\\s*(.*?)```"); + + private static final ObjectMapper MAPPER = new ObjectMapper() + .configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false); + + private JudgeJsonParser() { + } + + public static EvalJudgeService.JudgeResult parse(String text) { + if (text == null || text.isBlank()) { + throw new IllegalStateException("评委输出为空"); + } + String json = extractJson(text); + JudgeDto dto; + try { + dto = MAPPER.readValue(json, JudgeDto.class); + } catch (IOException e) { + throw new IllegalStateException("评委输出 JSON 解析失败: " + text, e); + } + if (dto.score == null) { + throw new IllegalStateException("评委输出缺少 score 字段: " + text); + } + double score = Math.max(0, Math.min(100, dto.score)); + boolean passed = Boolean.TRUE.equals(dto.passed); + List issues = dto.issues == null ? List.of() + : dto.issues.stream() + .map(i -> new EvalJudgeService.JudgeIssue(i.dimension, i.comment)) + .toList(); + return new EvalJudgeService.JudgeResult(score, passed, issues); + } + + private static String extractJson(String text) { + Matcher matcher = JSON_FENCE.matcher(text); + if (matcher.find()) { + return matcher.group(1).trim(); + } + int start = text.indexOf('{'); + int end = text.lastIndexOf('}'); + if (start < 0 || end < start) { + return text; + } + return text.substring(start, end + 1); + } + + private record JudgeDto(Double score, Boolean passed, List issues) { + } + + private record IssueDto(String dimension, String comment) { + } +} diff --git a/app/src/main/java/com/quantmore/modules/generator/eval/PythonSyntaxCheckService.java b/app/src/main/java/com/quantmore/modules/generator/eval/PythonSyntaxCheckService.java new file mode 100644 index 0000000..beab068 --- /dev/null +++ b/app/src/main/java/com/quantmore/modules/generator/eval/PythonSyntaxCheckService.java @@ -0,0 +1,141 @@ +package com.quantmore.modules.generator.eval; + +import lombok.extern.slf4j.Slf4j; +import org.springframework.stereotype.Component; + +import java.io.IOException; +import java.io.OutputStream; +import java.nio.charset.StandardCharsets; +import java.util.ArrayList; +import java.util.List; +import java.util.concurrent.TimeUnit; +import java.util.regex.Pattern; + +/** + * Python 语法检查:通过本机 python 子进程执行 ast.parse(代码经 stdin 传入)。 + * python 不可用时整体降级为 SKIPPED;另提供 Python 3.5 不兼容语法的启发式静态检查(仅警示)。 + */ +@Slf4j +@Component +public class PythonSyntaxCheckService { + + private static final Pattern PY_FSTRING = + Pattern.compile("(? warnings = detectPy35Incompatibilities(code); + try { + Process process = new ProcessBuilder(properties.getPythonBin(), "-c", PARSE_SCRIPT) + .start(); + try (OutputStream stdin = process.getOutputStream()) { + stdin.write(code.getBytes(StandardCharsets.UTF_8)); + } + if (!process.waitFor(properties.getPythonTimeout().toMillis(), TimeUnit.MILLISECONDS)) { + process.destroyForcibly(); + return new SyntaxCheckResult("FAIL", "语法检查超时", info.version(), warnings); + } + if (process.exitValue() == 0) { + return new SyntaxCheckResult("PASS", "", info.version(), warnings); + } + return new SyntaxCheckResult("FAIL", readStderr(process), info.version(), warnings); + } catch (IOException e) { + log.warn("语法检查进程异常: error={}", e.getMessage(), e); + return new SyntaxCheckResult("SKIPPED", "语法检查进程异常", info.version(), warnings); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + return new SyntaxCheckResult("SKIPPED", "语法检查被中断", info.version(), warnings); + } + } + + /** + * Python 3.5 不兼容语法启发式检查(正则误报率高,仅作警示项,不计入语法成败) + */ + public List detectPy35Incompatibilities(String code) { + List warnings = new ArrayList<>(); + if (PY_FSTRING.matcher(code).find()) { + warnings.add("f-string(Python 3.5 不支持)"); + } + if (PY_WALRUS.matcher(code).find()) { + warnings.add("海象运算符 :=(Python 3.8+)"); + } + if (PY_VAR_ANNOTATION.matcher(code).find()) { + warnings.add("变量注解(Python 3.6+)"); + } + if (PY_POSITIONAL_ONLY.matcher(code).find()) { + warnings.add("仅位置参数 /(Python 3.8+)"); + } + return warnings; + } + + private PythonInfo detectPython() { + try { + Process process = new ProcessBuilder(properties.getPythonBin(), "--version") + .redirectErrorStream(true) + .start(); + if (!process.waitFor(VERSION_TIMEOUT_SECONDS, TimeUnit.SECONDS)) { + process.destroyForcibly(); + return new PythonInfo(false, ""); + } + String version = new String(process.getInputStream().readAllBytes(), StandardCharsets.UTF_8) + .trim(); + return new PythonInfo(process.exitValue() == 0, version); + } catch (IOException e) { + return new PythonInfo(false, ""); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + return new PythonInfo(false, ""); + } + } + + private String readStderr(Process process) throws IOException { + String stderr = new String(process.getErrorStream().readAllBytes(), StandardCharsets.UTF_8) + .trim(); + return stderr.length() > STDERR_LIMIT ? stderr.substring(0, STDERR_LIMIT) : stderr; + } + + public record PythonInfo(boolean available, String version) { + } + + public record SyntaxCheckResult( + String status, + String message, + String pythonVersion, + List py35Warnings + ) { + } +} diff --git a/app/src/main/java/com/quantmore/modules/generator/eval/StrategyEvalRunner.java b/app/src/main/java/com/quantmore/modules/generator/eval/StrategyEvalRunner.java new file mode 100644 index 0000000..989f96c --- /dev/null +++ b/app/src/main/java/com/quantmore/modules/generator/eval/StrategyEvalRunner.java @@ -0,0 +1,59 @@ +package com.quantmore.modules.generator.eval; + +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; +import org.springframework.boot.CommandLineRunner; +import org.springframework.boot.SpringApplication; +import org.springframework.context.ConfigurableApplicationContext; +import org.springframework.core.annotation.Order; +import org.springframework.stereotype.Component; + +import java.util.List; + +/** + * 评测入口:APP_EVAL_ENABLED=true 时,启动后执行评测、写报告并退出进程。 + * 退出码:0=全部通过且知识库就绪;1=存在失败分支或知识库未就绪;2=致命错误(无 ADMIN/无 COMPLETED KB/用例加载失败)。 + * 顺序在 LocalKbSeedRunner(@Order(10)) 之后,且评测服务会等待向量化完成,双保险。 + */ +@Slf4j +@Component +@RequiredArgsConstructor +@Order(20) +public class StrategyEvalRunner implements CommandLineRunner { + + private final EvalProperties properties; + private final EvalCaseLoader caseLoader; + private final StrategyEvalService evalService; + private final EvalReportWriter reportWriter; + private final ConfigurableApplicationContext context; + + @Override + public void run(String... args) { + if (!properties.isEnabled()) { + return; + } + int result; + try { + log.info("策略生成评测开始: casesPath={}", properties.getCasesPath()); + List cases = caseLoader.load(); + EvalReport.FullReport report = evalService.run(cases); + reportWriter.write(report); + result = exitCodeOf(report); + log.info("策略生成评测结束: exitCode={}", result); + } catch (Exception e) { + log.error("策略生成评测致命失败", e); + result = 2; + } + final int exitCode = result; + int code = SpringApplication.exit(context, () -> exitCode); + System.exit(code); + } + + static int exitCodeOf(EvalReport.FullReport report) { + EvalReport.EvalSummary summary = report.summary(); + if (!report.kbReady() || summary.generationFailures() > 0 || summary.judgeFailures() > 0) { + return 1; + } + return 0; + } +} diff --git a/app/src/main/java/com/quantmore/modules/generator/eval/StrategyEvalService.java b/app/src/main/java/com/quantmore/modules/generator/eval/StrategyEvalService.java new file mode 100644 index 0000000..0adaece --- /dev/null +++ b/app/src/main/java/com/quantmore/modules/generator/eval/StrategyEvalService.java @@ -0,0 +1,193 @@ +package com.quantmore.modules.generator.eval; + +import com.quantmore.common.transaction.TransactionalExecutor; +import com.quantmore.modules.generator.dto.GenerateStrategyRequest; +import com.quantmore.modules.generator.dto.GenerateStrategyResponse; +import com.quantmore.modules.generator.repository.StrategyGenerationRepository; +import com.quantmore.modules.generator.service.StrategyGeneratorService; +import com.quantmore.modules.knowledgebase.model.VectorStatus; +import com.quantmore.modules.knowledgebase.repository.KnowledgeBaseRepository; +import com.quantmore.modules.user.model.UserEntity; +import com.quantmore.modules.user.model.UserPrincipal; +import com.quantmore.modules.user.model.UserRole; +import com.quantmore.modules.user.repository.UserRepository; +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; +import org.springframework.stereotype.Service; + +import java.time.LocalDateTime; +import java.util.ArrayList; +import java.util.List; +import java.util.concurrent.TimeUnit; + +/** + * 策略生成评测编排:每个用例跑 RAG / no-RAG 两分支,逐分支做语法检查与 LLM 评委评分, + * 最后聚合报告并清理本次评测产生的生成记录。不走 Controller,不受限流影响。 + */ +@Slf4j +@Service +@RequiredArgsConstructor +public class StrategyEvalService { + + private static final long KB_POLL_INTERVAL_MS = 2000; + + private final StrategyGeneratorService generatorService; + private final StrategyGenerationRepository generationRepository; + private final UserRepository userRepository; + private final KnowledgeBaseRepository knowledgeBaseRepository; + private final PythonSyntaxCheckService syntaxService; + private final EvalJudgeService judgeService; + private final EvalProperties properties; + private final TransactionalExecutor transactionalExecutor; + + public EvalReport.FullReport run(List cases) { + UserEntity admin = userRepository.findFirstByRoleOrderByIdAsc(UserRole.ADMIN) + .orElseThrow(() -> new IllegalStateException("评测需要至少一个 ADMIN 用户")); + UserPrincipal user = new UserPrincipal(admin.getId(), admin.getUsername(), admin.getRole()); + KbReadiness kb = waitForVectorReady(); + + List results = new ArrayList<>(); + List createdIds = new ArrayList<>(); + for (EvalCase caseMeta : cases) { + EvalReport.BranchResult rag = runBranch(caseMeta, user, true); + EvalReport.BranchResult noRag = runBranch(caseMeta, user, false); + if (rag.generationId() != null) { + createdIds.add(rag.generationId()); + } + if (noRag.generationId() != null) { + createdIds.add(noRag.generationId()); + } + results.add(new EvalReport.CaseResult(caseMeta, rag, noRag)); + log.info("评测用例完成: id={}, rag={}, noRag={}", + caseMeta.id(), branchBrief(rag), branchBrief(noRag)); + } + + cleanup(createdIds); + + String pythonVersion = syntaxService.pythonInfo().version(); + return new EvalReport.FullReport( + LocalDateTime.now().toString(), + admin.getUsername(), + resolveGenerateProvider(), + properties.getJudgeProvider(), + pythonVersion, + kb.ready(), + (int) kb.completed(), + (int) kb.failed(), + results, + EvalReport.EvalSummary.of(results, properties.getJudgePassScore()) + ); + } + + private EvalReport.BranchResult runBranch(EvalCase caseMeta, UserPrincipal user, boolean ragEnabled) { + long start = System.currentTimeMillis(); + try { + GenerateStrategyResponse response = + generatorService.generateForUser(toRequest(caseMeta, ragEnabled), user); + long elapsed = System.currentTimeMillis() - start; + + PythonSyntaxCheckService.SyntaxCheckResult syntax = syntaxService.check(response.code()); + EvalReport.BranchResult branch; + try { + EvalJudgeService.JudgeResult judge = judgeService.judge(caseMeta, response.code()); + branch = new EvalReport.BranchResult( + ragEnabled, true, null, response.id(), elapsed, syntax, true, judge, null); + } catch (Exception e) { + branch = new EvalReport.BranchResult( + ragEnabled, true, null, response.id(), elapsed, syntax, false, null, e.getMessage()); + } + return branch; + } catch (Exception e) { + log.warn("评测分支失败: case={}, rag={}, error={}", + caseMeta.id(), ragEnabled, e.getMessage(), e); + long elapsed = System.currentTimeMillis() - start; + return new EvalReport.BranchResult( + ragEnabled, false, e.getMessage(), null, elapsed, null, false, null, null); + } + } + + private GenerateStrategyRequest toRequest(EvalCase caseMeta, boolean ragEnabled) { + return new GenerateStrategyRequest( + caseMeta.name(), + caseMeta.market(), + caseMeta.frequency(), + caseMeta.buyConditions(), + caseMeta.sellConditions(), + caseMeta.riskControls(), + ragEnabled ? null : List.of(), + resolveGenerateProvider(), + ragEnabled ? null : true + ); + } + + private String resolveGenerateProvider() { + String provider = properties.getGenerateProvider(); + return (provider == null || provider.isBlank()) ? null : provider; + } + + /** + * 等待所有知识库向量化进入终态(PENDING/PROCESSING 清零或超时); + * 无任何 COMPLETED 知识库视为致命错误。 + */ + private KbReadiness waitForVectorReady() { + long deadline = System.currentTimeMillis() + properties.getVectorWaitTimeout().toMillis(); + boolean ready = false; + while (System.currentTimeMillis() < deadline) { + long pending = knowledgeBaseRepository.countByVectorStatus(VectorStatus.PENDING); + long processing = knowledgeBaseRepository.countByVectorStatus(VectorStatus.PROCESSING); + if (pending + processing == 0) { + ready = true; + break; + } + sleep(KB_POLL_INTERVAL_MS); + } + long completed = knowledgeBaseRepository.countByVectorStatus(VectorStatus.COMPLETED); + long failed = knowledgeBaseRepository.countByVectorStatus(VectorStatus.FAILED); + if (completed == 0) { + throw new IllegalStateException( + "无任何 COMPLETED 知识库,请先设置 APP_SEED_KB_DIR=docs 启动导入种子知识库"); + } + if (!ready) { + log.warn("等待向量化超时(仍有 PENDING/PROCESSING),评测继续,RAG 分支可能检索不到内容"); + } + if (failed > 0) { + knowledgeBaseRepository.findByVectorStatusOrderByUploadedAtDesc(VectorStatus.FAILED) + .forEach(kb -> log.warn("向量化 FAILED 知识库: id={}, name={}", kb.getId(), kb.getName())); + } + return new KbReadiness(ready, completed, failed); + } + + private void cleanup(List ids) { + if (ids.isEmpty() || !properties.isCleanupRecords()) { + return; + } + try { + transactionalExecutor.run(() -> generationRepository.deleteAllByIdInBatch(ids)); + log.info("评测生成记录已清理: count={}", ids.size()); + } catch (Exception e) { + log.warn("评测生成记录清理失败: count={}, error={}", ids.size(), e.getMessage(), e); + } + } + + private String branchBrief(EvalReport.BranchResult branch) { + if (!branch.generationOk()) { + return "生成失败"; + } + String syntax = branch.syntax() == null ? "-" : branch.syntax().status(); + if (!branch.judgeOk()) { + return "语法" + syntax + "/评委失败"; + } + return "语法" + syntax + "/评分" + branch.judge().score(); + } + + private void sleep(long millis) { + try { + TimeUnit.MILLISECONDS.sleep(millis); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + } + + private record KbReadiness(boolean ready, long completed, long failed) { + } +} diff --git a/app/src/main/java/com/quantmore/modules/generator/service/StrategyGeneratorService.java b/app/src/main/java/com/quantmore/modules/generator/service/StrategyGeneratorService.java index f8ddddf..8ccfc7b 100644 --- a/app/src/main/java/com/quantmore/modules/generator/service/StrategyGeneratorService.java +++ b/app/src/main/java/com/quantmore/modules/generator/service/StrategyGeneratorService.java @@ -4,6 +4,7 @@ import com.quantmore.common.ai.PromptSanitizer; import com.quantmore.common.exception.BusinessException; import com.quantmore.common.exception.ErrorCode; +import com.quantmore.common.transaction.TransactionalExecutor; import com.quantmore.modules.generator.config.GeneratorProperties; import com.quantmore.modules.generator.dto.GenerateStrategyRequest; import com.quantmore.modules.generator.dto.GenerateStrategyResponse; @@ -39,12 +40,16 @@ public class StrategyGeneratorService { static final Pattern PYTHON_FENCE = Pattern.compile("```python\\s*(.*?)```", Pattern.DOTALL); + static final String NO_RAG_CONTEXT = + "(跳过知识库检索:本次为无 RAG 对照生成,请严格按 PTrade 官方规范编写,不确定的 API 用注释标注)"; + private final StrategyGenerationRepository repository; private final KnowledgeBaseVectorService vectorService; private final KnowledgeBaseRepository knowledgeBaseRepository; private final LlmProviderRegistry registry; private final CurrentUserService currentUserService; private final PromptSanitizer sanitizer; + private final TransactionalExecutor transactionalExecutor; private final PromptTemplate systemPromptTemplate; private final PromptTemplate userPromptTemplate; private final int topK; @@ -58,13 +63,15 @@ public StrategyGeneratorService( CurrentUserService currentUserService, PromptSanitizer sanitizer, GeneratorProperties properties, - ResourceLoader resourceLoader) throws IOException { + ResourceLoader resourceLoader, + TransactionalExecutor transactionalExecutor) throws IOException { this.repository = repository; this.vectorService = vectorService; this.knowledgeBaseRepository = knowledgeBaseRepository; this.registry = registry; this.currentUserService = currentUserService; this.sanitizer = sanitizer; + this.transactionalExecutor = transactionalExecutor; this.systemPromptTemplate = new PromptTemplate( resourceLoader.getResource(properties.getSystemPromptPath()) .getContentAsString(StandardCharsets.UTF_8) @@ -80,10 +87,16 @@ public StrategyGeneratorService( /** * 生成策略 */ - @Transactional public GenerateStrategyResponse generate(GenerateStrategyRequest request) { UserPrincipal user = currentUserService.get(); + return generateForUser(request, user); + } + /** + * 生成策略(显式传入用户,供评测等无 SecurityContext 的场景使用) + * LLM 调用在事务外执行,仅持久化走小范围事务 + */ + public GenerateStrategyResponse generateForUser(GenerateStrategyRequest request, UserPrincipal user) { // 1. 输入清洗(反注入) String strategyName = sanitizer.sanitize(request.strategyName()).trim(); String buyConditions = sanitizer.sanitize(request.buyConditions()).trim(); @@ -92,20 +105,26 @@ public GenerateStrategyResponse generate(GenerateStrategyRequest request) { String riskControls = request.riskControls() == null ? "" : sanitizer.sanitize(request.riskControls()).trim(); - // 2. 解析知识库范围并校验可见性 - List kbIds = resolveKnowledgeBaseIds(user, request.knowledgeBaseIds()); - if (user.role() != UserRole.ADMIN) { - for (Long kbId : kbIds) { - if (!knowledgeBaseRepository.isVisibleToUser(kbId, user.id())) { - throw new BusinessException(ErrorCode.KNOWLEDGE_BASE_FORBIDDEN, - "知识库不可见或不存在: " + kbId); + // 2. 解析知识库范围并校验可见性(跳过检索时不解析、不校验) + List kbIds; + String context; + if (Boolean.TRUE.equals(request.skipRetrieval())) { + kbIds = List.of(); + context = NO_RAG_CONTEXT; + } else { + kbIds = resolveKnowledgeBaseIds(user, request.knowledgeBaseIds()); + if (user.role() != UserRole.ADMIN) { + for (Long kbId : kbIds) { + if (!knowledgeBaseRepository.isVisibleToUser(kbId, user.id())) { + throw new BusinessException(ErrorCode.KNOWLEDGE_BASE_FORBIDDEN, + "知识库不可见或不存在: " + kbId); + } } } + // 3. 检索参考示例(检索失败不阻断生成) + context = retrieveContext(strategyName, request, buyConditions, sellConditions, kbIds); } - // 3. 检索参考示例(检索失败不阻断生成) - String context = retrieveContext(strategyName, request, buyConditions, sellConditions, kbIds); - // 4. 渲染提示词并生成 String systemPrompt = systemPromptTemplate.render(Map.of("context", context)); String userPrompt = userPromptTemplate.render(Map.of( @@ -136,24 +155,26 @@ public GenerateStrategyResponse generate(GenerateStrategyRequest request) { SplitResult split = splitExplanationAndCode(raw); - // 5. 持久化 - StrategyGenerationEntity entity = StrategyGenerationEntity.builder() - .userId(user.id()) - .strategyName(strategyName) - .market(request.market()) - .frequency(request.frequency()) - .buyConditions(buyConditions) - .sellConditions(sellConditions) - .riskControls(riskControls) - .knowledgeBaseIds(kbIds.stream().map(String::valueOf).collect(Collectors.joining(","))) - .providerId(request.providerId()) - .generatedCode(split.code()) - .explanation(split.explanation()) - .build(); - entity = repository.save(entity); - - log.info("策略生成完成: id={}, strategy={}, userId={}", entity.getId(), strategyName, user.id()); - return toResponse(entity); + // 5. 持久化(唯一事务点,LLM 调用已在上方事务外完成) + GenerateStrategyResponse response = transactionalExecutor.call(() -> { + StrategyGenerationEntity entity = StrategyGenerationEntity.builder() + .userId(user.id()) + .strategyName(strategyName) + .market(request.market()) + .frequency(request.frequency()) + .buyConditions(buyConditions) + .sellConditions(sellConditions) + .riskControls(riskControls) + .knowledgeBaseIds(kbIds.stream().map(String::valueOf).collect(Collectors.joining(","))) + .providerId(request.providerId()) + .generatedCode(split.code()) + .explanation(split.explanation()) + .build(); + return toResponse(repository.save(entity)); + }); + + log.info("策略生成完成: id={}, strategy={}, userId={}", response.id(), strategyName, user.id()); + return response; } /** diff --git a/app/src/main/java/com/quantmore/modules/knowledgebase/seed/LocalKbSeedRunner.java b/app/src/main/java/com/quantmore/modules/knowledgebase/seed/LocalKbSeedRunner.java index b651941..a8e2c51 100644 --- a/app/src/main/java/com/quantmore/modules/knowledgebase/seed/LocalKbSeedRunner.java +++ b/app/src/main/java/com/quantmore/modules/knowledgebase/seed/LocalKbSeedRunner.java @@ -13,6 +13,7 @@ import lombok.extern.slf4j.Slf4j; import org.springframework.beans.factory.annotation.Value; import org.springframework.boot.CommandLineRunner; +import org.springframework.core.annotation.Order; import org.springframework.stereotype.Component; import org.springframework.web.multipart.MultipartFile; @@ -42,6 +43,7 @@ @Component @RequiredArgsConstructor @Slf4j +@Order(10) public class LocalKbSeedRunner implements CommandLineRunner { private final UserRepository userRepository; diff --git a/app/src/main/resources/application.yml b/app/src/main/resources/application.yml index c7da64c..5f2ea3a 100644 --- a/app/src/main/resources/application.yml +++ b/app/src/main/resources/application.yml @@ -101,6 +101,17 @@ app: user-prompt-path: ${APP_GENERATOR_USER_PROMPT_PATH:classpath:prompts/strategy-generator-user.st} top-k: ${APP_GENERATOR_TOPK:8} min-score: ${APP_GENERATOR_MIN_SCORE:0.18} + eval: + enabled: ${APP_EVAL_ENABLED:false} + cases-path: ${APP_EVAL_CASES_PATH:classpath:eval/strategy-eval-cases.json} + generate-provider: ${APP_EVAL_GENERATE_PROVIDER:} + judge-provider: ${APP_EVAL_JUDGE_PROVIDER:} + output-dir: ${APP_EVAL_OUTPUT_DIR:eval-output} + judge-pass-score: ${APP_EVAL_JUDGE_PASS_SCORE:70.0} + python-bin: ${APP_EVAL_PYTHON_BIN:python3} + python-timeout: ${APP_EVAL_PYTHON_TIMEOUT:10s} + vector-wait-timeout: ${APP_EVAL_VECTOR_WAIT_TIMEOUT:120s} + cleanup-records: ${APP_EVAL_CLEANUP_RECORDS:true} ai: default-provider: qwen default-embedding-provider: qwen diff --git a/app/src/main/resources/eval/strategy-eval-cases.json b/app/src/main/resources/eval/strategy-eval-cases.json new file mode 100644 index 0000000..438042a --- /dev/null +++ b/app/src/main/resources/eval/strategy-eval-cases.json @@ -0,0 +1,122 @@ +[ + { + "id": "s01", + "name": "双均线金叉死叉", + "market": "STOCK", + "frequency": "DAILY", + "buyConditions": "5日均线上穿10日均线时买入", + "sellConditions": "5日均线下穿10日均线时卖出", + "riskControls": "", + "difficulty": "SIMPLE" + }, + { + "id": "s02", + "name": "MACD金叉死叉", + "market": "STOCK", + "frequency": "DAILY", + "buyConditions": "MACD指标金叉(DIF上穿DEA)时买入", + "sellConditions": "MACD指标死叉(DIF下穿DEA)时卖出", + "riskControls": "", + "difficulty": "SIMPLE" + }, + { + "id": "s03", + "name": "双均线加止损止盈", + "market": "STOCK", + "frequency": "DAILY", + "buyConditions": "5日均线上穿10日均线时买入", + "sellConditions": "5日均线下穿10日均线时卖出", + "riskControls": "单只股票仓位不超过总资产50%;持仓亏损达5%止损;盈利达10%止盈", + "difficulty": "MEDIUM" + }, + { + "id": "s04", + "name": "集合竞价追涨停", + "market": "STOCK", + "frequency": "DAILY", + "buyConditions": "每天9:23集合竞价阶段,判断昨日涨停股今日竞价高开且接近涨停价时以涨停价挂单买入", + "sellConditions": "次日开盘后不涨停即卖出", + "riskControls": "单只股票仓位不超过总资产30%", + "difficulty": "MEDIUM" + }, + { + "id": "s05", + "name": "多因子轮动", + "market": "STOCK", + "frequency": "DAILY", + "buyConditions": "从沪深300指数成分股中选出20日动量排名前10且收盘价在60日均线上方的股票,等权重买入", + "sellConditions": "股票跌出动量前10名或跌破60日均线时卖出", + "riskControls": "停牌股票跳过;单只股票仓位不超过总资产20%;下单前检查未成交订单防止重复下单", + "difficulty": "COMPLEX" + }, + { + "id": "s06", + "name": "5分钟均线突破", + "market": "STOCK", + "frequency": "MINUTE", + "buyConditions": "5分钟K线收盘价上穿20周期均线时买入", + "sellConditions": "5分钟K线收盘价下穿20周期均线时卖出", + "riskControls": "当日累计亏损达3%后停止当日交易", + "difficulty": "MEDIUM" + }, + { + "id": "e01", + "name": "创业板ETF均线择时", + "market": "ETF", + "frequency": "DAILY", + "buyConditions": "创业板ETF收盘价站上20日均线时买入", + "sellConditions": "创业板ETF收盘价跌破20日均线时卖出", + "riskControls": "", + "difficulty": "SIMPLE" + }, + { + "id": "e02", + "name": "ETF分钟网格", + "market": "ETF", + "frequency": "MINUTE", + "buyConditions": "ETF价格跌破预设网格下限时买入一档", + "sellConditions": "ETF价格涨破预设网格上限时卖出一档", + "riskControls": "网格上下限价格差按ETF最小价差0.001元设置;仓位不超过总资产80%", + "difficulty": "MEDIUM" + }, + { + "id": "c01", + "name": "可转债双低轮动", + "market": "CONVERTIBLE_BOND", + "frequency": "DAILY", + "buyConditions": "每日收盘后按双低值(价格+溢价率)排序,买入排名前5的可转债", + "sellConditions": "可转债跌出双低排名前10时卖出", + "riskControls": "单只可转债仓位不超过总资产20%", + "difficulty": "MEDIUM" + }, + { + "id": "c02", + "name": "可转债轮动加强赎避让", + "market": "CONVERTIBLE_BOND", + "frequency": "DAILY", + "buyConditions": "每日收盘后按双低值排序买入排名前5的可转债,排除已公告强赎或临近到期的转债,溢价率超过30%的排除", + "sellConditions": "跌出双低排名前10、触发强赎公告或临近到期时卖出", + "riskControls": "单只可转债仓位不超过总资产20%", + "difficulty": "COMPLEX" + }, + { + "id": "f01", + "name": "股指期货均线趋势", + "market": "FUTURES", + "frequency": "DAILY", + "buyConditions": "沪深300股指期货IF主力合约收盘价站上20日均线时开多", + "sellConditions": "IF主力合约收盘价跌破20日均线时平多", + "riskControls": "单日最大亏损达2%时全部平仓;最大持仓手数不超过2手", + "difficulty": "MEDIUM" + }, + { + "id": "f02", + "name": "期货分钟动量突破", + "market": "FUTURES", + "frequency": "MINUTE", + "buyConditions": "1分钟K线收盘价突破前20周期最高价时开多,跌破前20周期最低价时开空", + "sellConditions": "多头持仓时价格回落至20周期均线下方平多;空头持仓时价格反弹至20周期均线上方平空", + "riskControls": "每天14:55强制平掉所有仓位;单日最大亏损达3%停止开仓;最大持仓手数不超过3手", + "difficulty": "COMPLEX" + } +] diff --git a/app/src/main/resources/prompts/strategy-eval-judge.st b/app/src/main/resources/prompts/strategy-eval-judge.st new file mode 100644 index 0000000..8e5d270 --- /dev/null +++ b/app/src/main/resources/prompts/strategy-eval-judge.st @@ -0,0 +1,29 @@ +# Role +你是 PTrade(恒生 PTrade 量化交易平台)资深策略评审专家,负责对生成的量化策略代码进行严格、客观的评分。 + +# 评分标准(总分 100) +1. **生命周期结构(15 分)**:`initialize` / `before_trading_start` / `handle_data`(或定时函数)是否齐全,职责划分是否正确。 +2. **买卖条件翻译(20 分)**:需求中的买入/卖出条件是否被精确、完整地翻译为代码逻辑,没有遗漏或篡改。 +3. **风控实现(15 分)**:止损/止盈/仓位上限等风控是否在交易逻辑**之前**判断;资金不足、持仓不足、可卖数量不足等防御是否到位。 +4. **API 使用合理性(20 分)**:只使用 PTrade 官方文档出现过的 API,不臆造;事件与 API 匹配(`get_snapshot` 不能用于回测;`get_price`/`get_trade_days` 的 `start_date` 与 `count` 不能同时传;`get_index_stocks` 不放 `initialize`);不确定的 API 应标注「需查阅文档」,未标注即扣分。 +5. **Python 3.5 兼容(10 分)**:禁止 f-string、海象运算符 `:=`、较新类型注解;禁止 `import os`。 +6. **可运行性(10 分)**:停牌、涨跌停、空行情等边界处理;下单前检查未成交订单防止重复下单;限价价格精度正确。 +7. **注释质量(10 分)**:关键段落与参数有中文注释。 + +# 待评审策略 +- 策略名称:{strategyName} +- 市场:{market} +- 频率:{frequency} +- 买入条件:{buyConditions} +- 卖出条件:{sellConditions} +- 风控要求:{riskControls} +- 生成代码: +{generatedCode} + +# 输出要求 +严格输出一个 JSON 对象,不要输出任何其他文字,不要使用代码围栏: +- 字段 score:0-100 的整数,按上述评分标准打分 +- 字段 passed:score 大于等于 70 时为 true,否则为 false +- 字段 issues:数组,每个元素包含 dimension(维度名字符串)与 comment(问题描述字符串)两个字段;只列扣分项,没有扣分项时为空数组 + +评分必须客观,宁严勿松。 diff --git a/app/src/test/java/com/quantmore/modules/generator/eval/EvalCaseLoaderTest.java b/app/src/test/java/com/quantmore/modules/generator/eval/EvalCaseLoaderTest.java new file mode 100644 index 0000000..e91ae35 --- /dev/null +++ b/app/src/test/java/com/quantmore/modules/generator/eval/EvalCaseLoaderTest.java @@ -0,0 +1,100 @@ +package com.quantmore.modules.generator.eval; + +import com.fasterxml.jackson.databind.ObjectMapper; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; +import org.springframework.core.io.DefaultResourceLoader; + +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +@DisplayName("EvalCaseLoader 测试") +class EvalCaseLoaderTest { + + @TempDir Path tempDir; + + private EvalProperties properties; + private EvalCaseLoader loader; + + @BeforeEach + void setUp() { + properties = new EvalProperties(); + loader = new EvalCaseLoader(new DefaultResourceLoader(), new ObjectMapper(), properties); + } + + private Path writeCases(String json) throws Exception { + Path file = tempDir.resolve("cases.json"); + Files.writeString(file, json); + properties.setCasesPath("file:" + file.toAbsolutePath()); + return file; + } + + @Test + @DisplayName("合法用例文件加载成功") + void loadsValidCases() throws Exception { + writeCases(""" + [ + {"id":"s01","name":"双均线","market":"STOCK","frequency":"DAILY", + "buyConditions":"五日均线上穿十日均线买入","sellConditions":"死叉卖出", + "riskControls":"","difficulty":"SIMPLE"} + ] + """); + + List cases = loader.load(); + + assertThat(cases).hasSize(1); + assertThat(cases.get(0).id()).isEqualTo("s01"); + assertThat(cases.get(0).market()).isEqualTo("STOCK"); + assertThat(cases.get(0).difficulty()).isEqualTo("SIMPLE"); + } + + @Test + @DisplayName("文件不存在抛 IllegalStateException") + void missingFileThrows() { + properties.setCasesPath("file:" + tempDir.resolve("nope.json").toAbsolutePath()); + + assertThatThrownBy(loader::load).isInstanceOf(IllegalStateException.class); + } + + @Test + @DisplayName("非法 market 枚举值抛 IllegalStateException 并携带用例 id") + void invalidEnumThrows() throws Exception { + writeCases(""" + [ + {"id":"b1","name":"x","market":"CRYPTO","frequency":"DAILY", + "buyConditions":"买入","sellConditions":"","riskControls":"","difficulty":"SIMPLE"} + ] + """); + + assertThatThrownBy(loader::load) + .isInstanceOf(IllegalStateException.class) + .hasMessageContaining("b1"); + } + + @Test + @DisplayName("空数组抛 IllegalStateException") + void emptyArrayThrows() throws Exception { + writeCases("[]"); + + assertThatThrownBy(loader::load).isInstanceOf(IllegalStateException.class); + } + + @Test + @DisplayName("必填字段为空抛 IllegalStateException") + void blankRequiredFieldThrows() throws Exception { + writeCases(""" + [ + {"id":"m1","name":"x","market":"STOCK","frequency":"DAILY", + "buyConditions":"","sellConditions":"","riskControls":"","difficulty":"SIMPLE"} + ] + """); + + assertThatThrownBy(loader::load).isInstanceOf(IllegalStateException.class); + } +} diff --git a/app/src/test/java/com/quantmore/modules/generator/eval/EvalJudgeServiceTest.java b/app/src/test/java/com/quantmore/modules/generator/eval/EvalJudgeServiceTest.java new file mode 100644 index 0000000..1e8196b --- /dev/null +++ b/app/src/test/java/com/quantmore/modules/generator/eval/EvalJudgeServiceTest.java @@ -0,0 +1,113 @@ +package com.quantmore.modules.generator.eval; + +import com.quantmore.common.ai.LlmProviderRegistry; +import com.quantmore.common.ai.PromptSanitizer; +import com.quantmore.common.ai.PromptSecurityConstants; +import com.quantmore.common.exception.BusinessException; +import com.quantmore.common.exception.ErrorCode; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.ai.chat.client.ChatClient; +import org.springframework.ai.chat.client.ChatClientResponse; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.ai.chat.model.Generation; + +import java.util.List; +import java.util.Map; +import java.util.concurrent.atomic.AtomicReference; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +@ExtendWith(MockitoExtension.class) +@DisplayName("EvalJudgeService 测试") +class EvalJudgeServiceTest { + + @Mock private LlmProviderRegistry registry; + @Mock private PromptSanitizer sanitizer; + + private EvalProperties properties; + private EvalJudgeService service; + + private static final String RAW_JSON = + "{\"score\": 82, \"passed\": true, " + + "\"issues\": [{\"dimension\": \"结构\", \"comment\": \"完整\"}]}"; + + @BeforeEach + void setUp() throws Exception { + properties = new EvalProperties(); + org.mockito.Mockito.lenient().when(sanitizer.sanitize(anyString())) + .thenAnswer(inv -> inv.getArgument(0)); + org.mockito.Mockito.lenient().when(sanitizer.wrapWithDelimiters(anyString(), anyString())) + .thenAnswer(inv -> "[" + inv.getArgument(0) + "]\n" + inv.getArgument(1)); + service = new EvalJudgeService(registry, sanitizer, properties); + } + + private AtomicReference stubPlainClient(String raw) { + ChatClient chatClient = mock(ChatClient.class); + ChatClient.ChatClientRequestSpec promptSpec = mock(ChatClient.ChatClientRequestSpec.class); + when(registry.getPlainChatClient()).thenReturn(chatClient); + when(chatClient.prompt()).thenReturn(promptSpec); + AtomicReference systemPrompt = new AtomicReference<>(); + when(promptSpec.system(anyString())).thenAnswer(inv -> { + systemPrompt.set(inv.getArgument(0)); + return promptSpec; + }); + when(promptSpec.call()).thenReturn(mock(ChatClient.CallResponseSpec.class)); + when(promptSpec.call().chatClientResponse()).thenReturn( + new ChatClientResponse( + new ChatResponse(List.of(new Generation(new AssistantMessage(raw)))), Map.of())); + return systemPrompt; + } + + private EvalCase caseMeta() { + return new EvalCase("s01", "双均线策略", "STOCK", "DAILY", "金叉买入", "死叉卖出", "", "SIMPLE"); + } + + @Test + @DisplayName("评分成功:提示词含用例字段/代码/反注入指令,结果解析正确") + void judgesAndParsesResult() { + AtomicReference systemPrompt = stubPlainClient(RAW_JSON); + + EvalJudgeService.JudgeResult result = service.judge(caseMeta(), "def initialize(context):\n pass\n"); + + assertThat(result.score()).isEqualTo(82.0); + assertThat(result.passed()).isTrue(); + assertThat(result.issues()).hasSize(1); + assertThat(systemPrompt.get()) + .contains("双均线策略", "金叉买入", "def initialize", "generated-code") + .endsWith(PromptSecurityConstants.ANTI_INJECTION_INSTRUCTION); + } + + @Test + @DisplayName("评委返回空内容抛 AI_SERVICE_ERROR") + void emptyRawThrowsAiServiceError() { + stubPlainClient(" "); + + assertThatThrownBy(() -> service.judge(caseMeta(), "pass\n")) + .isInstanceOf(BusinessException.class) + .satisfies(e -> assertThat(((BusinessException) e).getCode()) + .isEqualTo(ErrorCode.AI_SERVICE_ERROR.getCode())); + } + + @Test + @DisplayName("LLM 调用异常抛 AI_SERVICE_ERROR") + void llmFailureThrowsAiServiceError() { + ChatClient chatClient = mock(ChatClient.class); + when(registry.getPlainChatClient()).thenReturn(chatClient); + when(chatClient.prompt()).thenThrow(new RuntimeException("boom")); + + assertThatThrownBy(() -> service.judge(caseMeta(), "pass\n")) + .isInstanceOf(BusinessException.class) + .satisfies(e -> assertThat(((BusinessException) e).getCode()) + .isEqualTo(ErrorCode.AI_SERVICE_ERROR.getCode())); + } +} diff --git a/app/src/test/java/com/quantmore/modules/generator/eval/EvalReportWriterTest.java b/app/src/test/java/com/quantmore/modules/generator/eval/EvalReportWriterTest.java new file mode 100644 index 0000000..9fa31da --- /dev/null +++ b/app/src/test/java/com/quantmore/modules/generator/eval/EvalReportWriterTest.java @@ -0,0 +1,72 @@ +package com.quantmore.modules.generator.eval; + +import com.fasterxml.jackson.databind.ObjectMapper; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; + +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; + +import static org.assertj.core.api.Assertions.assertThat; + +@DisplayName("EvalReportWriter 测试") +class EvalReportWriterTest { + + @TempDir Path tempDir; + + private Path outputDir; + private EvalReportWriter writer; + + @BeforeEach + void setUp() { + outputDir = tempDir.resolve("out"); + EvalProperties properties = new EvalProperties(); + properties.setOutputDir(outputDir.toString()); + writer = new EvalReportWriter(properties, new ObjectMapper()); + } + + private EvalReport.BranchResult branch( + boolean ragEnabled, String syntaxStatus, boolean judgeOk, double score) { + PythonSyntaxCheckService.SyntaxCheckResult syntax = + new PythonSyntaxCheckService.SyntaxCheckResult(syntaxStatus, "", "3.9", List.of()); + EvalJudgeService.JudgeResult judge = new EvalJudgeService.JudgeResult(score, score >= 70, List.of()); + return new EvalReport.BranchResult( + ragEnabled, true, null, 1L, 100L, + syntax, judgeOk, judgeOk ? judge : null, judgeOk ? null : "评委失败"); + } + + private EvalReport.FullReport sampleReport() { + List results = List.of( + new EvalReport.CaseResult( + new EvalCase("s01", "双均线", "STOCK", "DAILY", "买入", "", "", "SIMPLE"), + branch(true, "PASS", true, 85), + branch(false, "PASS", true, 60)), + new EvalReport.CaseResult( + new EvalCase("s02", "轮动", "STOCK", "DAILY", "买入", "", "", "COMPLEX"), + branch(true, "FAIL", false, 0), + branch(false, "SKIPPED", true, 55))); + return new EvalReport.FullReport( + "2026-08-31T10:00:00", "admin", "qwen", "qwen", "Python 3.9.6", + true, 1, 0, results, EvalReport.EvalSummary.of(results, 70.0)); + } + + @Test + @DisplayName("写出 eval-report.json 与 eval-report.md,JSON 可反序列化回完整报告") + void writesJsonAndMarkdown() throws Exception { + writer.write(sampleReport()); + + Path json = outputDir.resolve("eval-report.json"); + Path md = outputDir.resolve("eval-report.md"); + assertThat(json).exists(); + assertThat(md).exists(); + + EvalReport.FullReport parsed = new ObjectMapper() + .readValue(Files.readString(json), EvalReport.FullReport.class); + assertThat(parsed.summary().totalCases()).isEqualTo(2); + assertThat(parsed.results()).hasSize(2); + assertThat(Files.readString(md)).contains("RAG", "no-RAG", "s01", "s02"); + } +} diff --git a/app/src/test/java/com/quantmore/modules/generator/eval/EvalSummaryTest.java b/app/src/test/java/com/quantmore/modules/generator/eval/EvalSummaryTest.java new file mode 100644 index 0000000..1b7a227 --- /dev/null +++ b/app/src/test/java/com/quantmore/modules/generator/eval/EvalSummaryTest.java @@ -0,0 +1,65 @@ +package com.quantmore.modules.generator.eval; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import java.util.List; + +import static org.assertj.core.api.Assertions.assertThat; + +@DisplayName("EvalSummary 聚合测试") +class EvalSummaryTest { + + private EvalCase caseMeta(String id) { + return new EvalCase(id, id, "STOCK", "DAILY", "买入", "", "", "SIMPLE"); + } + + private EvalReport.BranchResult branch( + boolean ragEnabled, boolean generationOk, String syntaxStatus, + boolean judgeOk, double score, List warnings) { + PythonSyntaxCheckService.SyntaxCheckResult syntax = + new PythonSyntaxCheckService.SyntaxCheckResult(syntaxStatus, "", "3.9", warnings); + EvalJudgeService.JudgeResult judge = new EvalJudgeService.JudgeResult(score, score >= 70, List.of()); + return new EvalReport.BranchResult( + ragEnabled, generationOk, generationOk ? null : "生成失败", 1L, 100L, + syntax, judgeOk, judgeOk ? judge : null, judgeOk ? "{}" : null); + } + + @Test + @DisplayName("聚合通过数/语法通过数/平均分/差值/失败计数") + void aggregatesPassedAndScores() { + EvalReport.CaseResult case1 = new EvalReport.CaseResult(caseMeta("s01"), + branch(true, true, "PASS", true, 85, List.of()), + branch(false, true, "PASS", true, 60, List.of())); + EvalReport.CaseResult case2 = new EvalReport.CaseResult(caseMeta("s02"), + branch(true, true, "FAIL", false, 0, List.of("f-string")), + branch(false, false, "SKIPPED", false, 0, List.of())); + + EvalReport.EvalSummary summary = EvalReport.EvalSummary.of(List.of(case1, case2), 70.0); + + assertThat(summary.totalCases()).isEqualTo(2); + assertThat(summary.ragPassed()).isEqualTo(1); + assertThat(summary.noRagPassed()).isEqualTo(0); + assertThat(summary.ragSyntaxPassed()).isEqualTo(1); + assertThat(summary.noRagSyntaxPassed()).isEqualTo(1); + assertThat(summary.generationFailures()).isEqualTo(1); + assertThat(summary.judgeFailures()).isEqualTo(1); + assertThat(summary.py35WarningCount()).isEqualTo(1); + assertThat(summary.ragAvgScore()).isEqualTo(85.0); + assertThat(summary.noRagAvgScore()).isEqualTo(60.0); + assertThat(summary.scoreDelta()).isEqualTo(25.0); + } + + @Test + @DisplayName("无评委结果时平均分为 0") + void noJudgeResultsMeansZeroAvg() { + EvalReport.CaseResult onlyFailure = new EvalReport.CaseResult(caseMeta("s03"), + branch(true, false, "SKIPPED", false, 0, List.of()), + branch(false, true, "FAIL", false, 0, List.of())); + + EvalReport.EvalSummary summary = EvalReport.EvalSummary.of(List.of(onlyFailure), 70.0); + + assertThat(summary.ragAvgScore()).isEqualTo(0.0); + assertThat(summary.noRagAvgScore()).isEqualTo(0.0); + } +} diff --git a/app/src/test/java/com/quantmore/modules/generator/eval/JudgeJsonParserTest.java b/app/src/test/java/com/quantmore/modules/generator/eval/JudgeJsonParserTest.java new file mode 100644 index 0000000..b33ffb3 --- /dev/null +++ b/app/src/test/java/com/quantmore/modules/generator/eval/JudgeJsonParserTest.java @@ -0,0 +1,65 @@ +package com.quantmore.modules.generator.eval; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +@DisplayName("JudgeJsonParser 测试") +class JudgeJsonParserTest { + + @Test + @DisplayName("解析 ```json 围栏内容") + void parsesFencedJson() { + EvalJudgeService.JudgeResult result = JudgeJsonParser.parse( + "以下是评分\n```json\n{\"score\": 85, \"passed\": true, " + + "\"issues\": [{\"dimension\": \"结构\", \"comment\": \"ok\"}]}\n```\n完毕"); + + assertThat(result.score()).isEqualTo(85.0); + assertThat(result.passed()).isTrue(); + assertThat(result.issues()).hasSize(1); + assertThat(result.issues().get(0).dimension()).isEqualTo("结构"); + } + + @Test + @DisplayName("解析裸 JSON(前后有杂文)") + void parsesBareJson() { + EvalJudgeService.JudgeResult result = JudgeJsonParser.parse( + "评分如下 {\"score\": 60, \"passed\": false, \"issues\": []} 谢谢"); + + assertThat(result.score()).isEqualTo(60.0); + assertThat(result.passed()).isFalse(); + } + + @Test + @DisplayName("issues 缺失时为空列表") + void missingIssuesDefaultsToEmpty() { + EvalJudgeService.JudgeResult result = JudgeJsonParser.parse("{\"score\": 70, \"passed\": true}"); + + assertThat(result.issues()).isEmpty(); + } + + @Test + @DisplayName("score 越界夹取到 [0,100]") + void clampsScore() { + assertThat(JudgeJsonParser.parse("{\"score\": 150, \"passed\": true}").score()) + .isEqualTo(100.0); + assertThat(JudgeJsonParser.parse("{\"score\": -5, \"passed\": false}").score()) + .isEqualTo(0.0); + } + + @Test + @DisplayName("非法 JSON 抛 IllegalStateException") + void invalidJsonThrows() { + assertThatThrownBy(() -> JudgeJsonParser.parse("不是 JSON")) + .isInstanceOf(IllegalStateException.class); + } + + @Test + @DisplayName("空输出抛 IllegalStateException") + void blankTextThrows() { + assertThatThrownBy(() -> JudgeJsonParser.parse(" ")) + .isInstanceOf(IllegalStateException.class); + } +} diff --git a/app/src/test/java/com/quantmore/modules/generator/eval/PythonSyntaxCheckServiceTest.java b/app/src/test/java/com/quantmore/modules/generator/eval/PythonSyntaxCheckServiceTest.java new file mode 100644 index 0000000..1f5a6cb --- /dev/null +++ b/app/src/test/java/com/quantmore/modules/generator/eval/PythonSyntaxCheckServiceTest.java @@ -0,0 +1,87 @@ +package com.quantmore.modules.generator.eval; + +import org.junit.jupiter.api.Assumptions; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; + +import java.nio.file.Path; +import java.time.Duration; +import java.util.List; + +import static org.assertj.core.api.Assertions.assertThat; + +@DisplayName("PythonSyntaxCheckService 测试") +class PythonSyntaxCheckServiceTest { + + @TempDir Path tempDir; + + private EvalProperties properties; + + @BeforeEach + void setUp() { + properties = new EvalProperties(); + properties.setPythonTimeout(Duration.ofSeconds(5)); + } + + private PythonSyntaxCheckService service() { + return new PythonSyntaxCheckService(properties); + } + + @Test + @DisplayName("合法 Python 代码语法检查 PASS 并携带 python 版本") + void validCodePasses() { + PythonSyntaxCheckService service = service(); + Assumptions.assumeTrue(service.pythonInfo().available(), "本机无 python3,跳过"); + + PythonSyntaxCheckService.SyntaxCheckResult result = + service.check("def initialize(context):\n pass\n"); + + assertThat(result.status()).isEqualTo("PASS"); + assertThat(result.pythonVersion()).isNotBlank(); + } + + @Test + @DisplayName("非法 Python 代码语法检查 FAIL 并携带错误信息") + void invalidCodeFails() { + PythonSyntaxCheckService service = service(); + Assumptions.assumeTrue(service.pythonInfo().available(), "本机无 python3,跳过"); + + PythonSyntaxCheckService.SyntaxCheckResult result = service.check("def initialize(:\n"); + + assertThat(result.status()).isEqualTo("FAIL"); + assertThat(result.message()).isNotBlank(); + } + + @Test + @DisplayName("python 解释器不存在时降级 SKIPPED") + void missingPythonSkips() { + properties.setPythonBin(tempDir.resolve("no-such-python").toString()); + + PythonSyntaxCheckService.SyntaxCheckResult result = service().check("anything"); + + assertThat(result.status()).isEqualTo("SKIPPED"); + } + + @Test + @DisplayName("检测 Python 3.5 不兼容语法启发式") + void detectsPy35Incompatibilities() { + List warnings = service().detectPy35Incompatibilities( + "s = f'当前价格 {price}'\n" + + "if (y := 1):\n" + + " pass\n"); + + assertThat(warnings).anyMatch(w -> w.contains("f-string")); + assertThat(warnings).anyMatch(w -> w.contains("海象")); + } + + @Test + @DisplayName("普通代码无 3.5 不兼容警示") + void normalCodeHasNoWarnings() { + List warnings = service().detectPy35Incompatibilities( + "def handle_data(context, data):\n return\n"); + + assertThat(warnings).isEmpty(); + } +} diff --git a/app/src/test/java/com/quantmore/modules/generator/eval/StrategyEvalRunnerTest.java b/app/src/test/java/com/quantmore/modules/generator/eval/StrategyEvalRunnerTest.java new file mode 100644 index 0000000..ce2ef30 --- /dev/null +++ b/app/src/test/java/com/quantmore/modules/generator/eval/StrategyEvalRunnerTest.java @@ -0,0 +1,43 @@ +package com.quantmore.modules.generator.eval; + +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import java.util.List; + +import static org.assertj.core.api.Assertions.assertThat; + +@DisplayName("StrategyEvalRunner 退出码测试") +class StrategyEvalRunnerTest { + + private EvalReport.FullReport report(boolean kbReady, int generationFailures, int judgeFailures) { + EvalReport.EvalSummary summary = new EvalReport.EvalSummary( + 1, 1, 1, 80, 60, 1, 1, generationFailures, judgeFailures, 0, 20); + return new EvalReport.FullReport( + "now", "admin", null, null, "3.9", kbReady, 1, 0, List.of(), summary); + } + + @Test + @DisplayName("知识库就绪且无失败时退出码 0") + void allPassedExitsZero() { + assertThat(StrategyEvalRunner.exitCodeOf(report(true, 0, 0))).isEqualTo(0); + } + + @Test + @DisplayName("存在生成失败时退出码 1") + void generationFailureExitsOne() { + assertThat(StrategyEvalRunner.exitCodeOf(report(true, 1, 0))).isEqualTo(1); + } + + @Test + @DisplayName("存在评委失败时退出码 1") + void judgeFailureExitsOne() { + assertThat(StrategyEvalRunner.exitCodeOf(report(true, 0, 1))).isEqualTo(1); + } + + @Test + @DisplayName("知识库未就绪时退出码 1") + void kbNotReadyExitsOne() { + assertThat(StrategyEvalRunner.exitCodeOf(report(false, 0, 0))).isEqualTo(1); + } +} diff --git a/app/src/test/java/com/quantmore/modules/generator/eval/StrategyEvalServiceTest.java b/app/src/test/java/com/quantmore/modules/generator/eval/StrategyEvalServiceTest.java new file mode 100644 index 0000000..3e3cabb --- /dev/null +++ b/app/src/test/java/com/quantmore/modules/generator/eval/StrategyEvalServiceTest.java @@ -0,0 +1,182 @@ +package com.quantmore.modules.generator.eval; + +import com.quantmore.common.exception.BusinessException; +import com.quantmore.common.exception.ErrorCode; +import com.quantmore.common.transaction.TransactionalExecutor; +import com.quantmore.modules.generator.dto.GenerateStrategyRequest; +import com.quantmore.modules.generator.dto.GenerateStrategyResponse; +import com.quantmore.modules.generator.repository.StrategyGenerationRepository; +import com.quantmore.modules.generator.service.StrategyGeneratorService; +import com.quantmore.modules.knowledgebase.model.VectorStatus; +import com.quantmore.modules.knowledgebase.repository.KnowledgeBaseRepository; +import com.quantmore.modules.user.model.UserEntity; +import com.quantmore.modules.user.model.UserRole; +import com.quantmore.modules.user.repository.UserRepository; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import java.time.Duration; +import java.time.LocalDateTime; +import java.util.List; +import java.util.Optional; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +@ExtendWith(MockitoExtension.class) +@DisplayName("StrategyEvalService 测试") +class StrategyEvalServiceTest { + + @Mock private StrategyGeneratorService generatorService; + @Mock private StrategyGenerationRepository generationRepository; + @Mock private UserRepository userRepository; + @Mock private KnowledgeBaseRepository knowledgeBaseRepository; + @Mock private PythonSyntaxCheckService syntaxService; + @Mock private EvalJudgeService judgeService; + @Mock private TransactionalExecutor transactionalExecutor; + + private EvalProperties properties; + private StrategyEvalService service; + + @BeforeEach + void setUp() { + properties = new EvalProperties(); + properties.setVectorWaitTimeout(Duration.ofSeconds(1)); + service = new StrategyEvalService( + generatorService, generationRepository, userRepository, knowledgeBaseRepository, + syntaxService, judgeService, properties, transactionalExecutor); + } + + private EvalCase caseMeta(String id) { + return new EvalCase(id, "策略" + id, "STOCK", "DAILY", "金叉买入", "死叉卖出", "", "SIMPLE"); + } + + private UserEntity admin() { + UserEntity admin = new UserEntity(); + admin.setId(1L); + admin.setUsername("admin"); + admin.setRole(UserRole.ADMIN); + return admin; + } + + private void stubReadyKb() { + when(userRepository.findFirstByRoleOrderByIdAsc(UserRole.ADMIN)) + .thenReturn(Optional.of(admin())); + when(knowledgeBaseRepository.countByVectorStatus(VectorStatus.PENDING)).thenReturn(0L); + when(knowledgeBaseRepository.countByVectorStatus(VectorStatus.PROCESSING)).thenReturn(0L); + when(knowledgeBaseRepository.countByVectorStatus(VectorStatus.COMPLETED)).thenReturn(1L); + when(knowledgeBaseRepository.countByVectorStatus(VectorStatus.FAILED)).thenReturn(0L); + } + + private GenerateStrategyResponse response() { + return new GenerateStrategyResponse( + 1L, "策略", "策略.py", "STOCK", "DAILY", "def initialize(context):\n pass\n", + "说明", null, LocalDateTime.now()); + } + + @Test + @DisplayName("每个用例跑 RAG 与 no-RAG 两分支,skipRetrieval 分别为 null/true,共清理 2 条记录") + void runsBothBranchesPerCase() { + stubReadyKb(); + when(syntaxService.pythonInfo()) + .thenReturn(new PythonSyntaxCheckService.PythonInfo(true, "Python 3.9.6")); + when(syntaxService.check(anyString())) + .thenReturn(new PythonSyntaxCheckService.SyntaxCheckResult("PASS", "", "3.9", List.of())); + when(judgeService.judge(any(), anyString())) + .thenReturn(new EvalJudgeService.JudgeResult(85, true, List.of())); + when(generatorService.generateForUser(any(), any())).thenReturn(response()); + + EvalReport.FullReport report = service.run(List.of(caseMeta("s01"))); + + assertThat(report.results()).hasSize(1); + assertThat(report.kbReady()).isTrue(); + assertThat(report.kbCompleted()).isEqualTo(1); + assertThat(report.summary().ragPassed()).isEqualTo(1); + assertThat(report.summary().noRagPassed()).isEqualTo(1); + + ArgumentCaptor reqCaptor = + ArgumentCaptor.forClass(GenerateStrategyRequest.class); + verify(generatorService, times(2)).generateForUser(reqCaptor.capture(), any()); + assertThat(reqCaptor.getAllValues().get(0).skipRetrieval()).isNull(); + assertThat(reqCaptor.getAllValues().get(1).skipRetrieval()).isTrue(); + + ArgumentCaptor runnableCaptor = ArgumentCaptor.forClass(Runnable.class); + verify(transactionalExecutor).run(runnableCaptor.capture()); + runnableCaptor.getValue().run(); + verify(generationRepository).deleteAllByIdInBatch(List.of(1L, 1L)); + } + + @Test + @DisplayName("分支生成异常记为 generationFailed,不中断后续用例,该分支不做语法与评委") + void generationFailureDoesNotInterrupt() { + stubReadyKb(); + when(syntaxService.pythonInfo()) + .thenReturn(new PythonSyntaxCheckService.PythonInfo(true, "Python 3.9.6")); + when(syntaxService.check(anyString())) + .thenReturn(new PythonSyntaxCheckService.SyntaxCheckResult("PASS", "", "3.9", List.of())); + when(judgeService.judge(any(), anyString())) + .thenReturn(new EvalJudgeService.JudgeResult(80, true, List.of())); + when(generatorService.generateForUser(any(), any())) + .thenThrow(new BusinessException(ErrorCode.AI_SERVICE_ERROR, "boom")) + .thenReturn(response()); + + EvalReport.FullReport report = service.run(List.of(caseMeta("s01"))); + + assertThat(report.summary().generationFailures()).isEqualTo(1); + assertThat(report.summary().judgeFailures()).isEqualTo(0); + verify(syntaxService, times(1)).check(anyString()); + verify(judgeService, times(1)).judge(any(), anyString()); + } + + @Test + @DisplayName("评委失败记为 judgeFailed,不影响统计完成") + void judgeFailureIsRecorded() { + stubReadyKb(); + when(syntaxService.pythonInfo()) + .thenReturn(new PythonSyntaxCheckService.PythonInfo(true, "Python 3.9.6")); + when(syntaxService.check(anyString())) + .thenReturn(new PythonSyntaxCheckService.SyntaxCheckResult("PASS", "", "3.9", List.of())); + when(judgeService.judge(any(), anyString())) + .thenThrow(new BusinessException(ErrorCode.AI_SERVICE_ERROR, "评委失败")); + when(generatorService.generateForUser(any(), any())).thenReturn(response()); + + EvalReport.FullReport report = service.run(List.of(caseMeta("s01"))); + + assertThat(report.summary().judgeFailures()).isEqualTo(2); + assertThat(report.summary().ragPassed()).isEqualTo(0); + } + + @Test + @DisplayName("无 ADMIN 用户抛 IllegalStateException") + void noAdminThrows() { + when(userRepository.findFirstByRoleOrderByIdAsc(UserRole.ADMIN)).thenReturn(Optional.empty()); + + assertThatThrownBy(() -> service.run(List.of(caseMeta("s01")))) + .isInstanceOf(IllegalStateException.class) + .hasMessageContaining("ADMIN"); + } + + @Test + @DisplayName("无 COMPLETED 知识库抛 IllegalStateException") + void noCompletedKbThrows() { + when(userRepository.findFirstByRoleOrderByIdAsc(UserRole.ADMIN)) + .thenReturn(Optional.of(admin())); + when(knowledgeBaseRepository.countByVectorStatus(VectorStatus.PENDING)).thenReturn(0L); + when(knowledgeBaseRepository.countByVectorStatus(VectorStatus.PROCESSING)).thenReturn(0L); + when(knowledgeBaseRepository.countByVectorStatus(VectorStatus.COMPLETED)).thenReturn(0L); + + assertThatThrownBy(() -> service.run(List.of(caseMeta("s01")))) + .isInstanceOf(IllegalStateException.class) + .hasMessageContaining("COMPLETED"); + } +} diff --git a/app/src/test/java/com/quantmore/modules/generator/service/StrategyGeneratorServiceTest.java b/app/src/test/java/com/quantmore/modules/generator/service/StrategyGeneratorServiceTest.java index b64b002..0a00839 100644 --- a/app/src/test/java/com/quantmore/modules/generator/service/StrategyGeneratorServiceTest.java +++ b/app/src/test/java/com/quantmore/modules/generator/service/StrategyGeneratorServiceTest.java @@ -56,6 +56,7 @@ class StrategyGeneratorServiceTest { @Mock private CurrentUserService currentUserService; @Mock private PromptSanitizer sanitizer; @Mock private ResourceLoader resourceLoader; + @Mock private com.quantmore.common.transaction.TransactionalExecutor transactionalExecutor; private StrategyGeneratorService service; @@ -65,7 +66,7 @@ private GenerateStrategyRequest request() { return new GenerateStrategyRequest( "双均线策略", "STOCK", "DAILY", "五日均线上穿十日均线买入", "五日均线下穿十日均线卖出", - "单只股票仓位不超过总资产50%", List.of(1L), null); + "单只股票仓位不超过总资产50%", List.of(1L), null, null); } @BeforeEach @@ -85,7 +86,32 @@ void setUp() throws Exception { service = new StrategyGeneratorService( repository, vectorService, knowledgeBaseRepository, registry, - currentUserService, sanitizer, properties, resourceLoader); + currentUserService, sanitizer, properties, resourceLoader, transactionalExecutor); + org.mockito.Mockito.lenient().when(transactionalExecutor.call(any())) + .thenAnswer(inv -> ((java.util.function.Supplier) inv.getArgument(0)).get()); + } + + private java.util.concurrent.atomic.AtomicReference stubChatAndSave() { + ChatClient chatClient = mock(ChatClient.class); + ChatClient.ChatClientRequestSpec promptSpec = mock(ChatClient.ChatClientRequestSpec.class); + when(registry.getChatClientForUser(any(), any())).thenReturn(chatClient); + when(chatClient.prompt()).thenReturn(promptSpec); + java.util.concurrent.atomic.AtomicReference systemPrompt = + new java.util.concurrent.atomic.AtomicReference<>(); + when(promptSpec.system(anyString())).thenAnswer(inv -> { + systemPrompt.set(inv.getArgument(0)); + return promptSpec; + }); + when(promptSpec.user(anyString())).thenReturn(promptSpec); + when(promptSpec.call()).thenReturn(mock(ChatClient.CallResponseSpec.class)); + when(promptSpec.call().content()).thenReturn("```python\npass\n```"); + when(repository.save(any())).thenAnswer(inv -> { + StrategyGenerationEntity e = inv.getArgument(0); + e.setId(1L); + e.setCreatedAt(java.time.LocalDateTime.now()); + return e; + }); + return systemPrompt; } @Nested @@ -192,11 +218,72 @@ void generatesWithoutAnyKb() { }); GenerateStrategyResponse response = service.generate(new GenerateStrategyRequest( - "测试策略", "STOCK", "DAILY", "买入", "", "", null, null)); + "测试策略", "STOCK", "DAILY", "买入", "", "", null, null, null)); + + assertThat(response).isNotNull(); + verify(vectorService, org.mockito.Mockito.never()) + .similaritySearch(anyString(), anyList(), anyInt(), anyDouble()); + } + } + + @Nested + @DisplayName("跳过检索") + class SkipRetrieval { + + private GenerateStrategyRequest skipRequest() { + return new GenerateStrategyRequest( + "无 RAG 对照策略", "STOCK", "DAILY", + "五日均线上穿十日均线买入", "五日均线下穿十日均线卖出", + "单只股票仓位不超过总资产50%", List.of(1L), null, true); + } + + @Test + @DisplayName("skipRetrieval=true 时不检索、不解析知识库范围,提示词携带无 RAG 说明,落库 knowledgeBaseIds 为空") + void skipsRetrievalAndKbResolution() { + when(currentUserService.get()).thenReturn(new UserPrincipal(1L, "admin", UserRole.ADMIN)); + java.util.concurrent.atomic.AtomicReference systemPrompt = stubChatAndSave(); + + GenerateStrategyResponse response = service.generate(skipRequest()); assertThat(response).isNotNull(); verify(vectorService, org.mockito.Mockito.never()) .similaritySearch(anyString(), anyList(), anyInt(), anyDouble()); + verify(knowledgeBaseRepository, org.mockito.Mockito.never()) + .findAllByOrderByUploadedAtDesc(); + assertThat(systemPrompt.get()).isNotNull().contains("跳过知识库检索"); + ArgumentCaptor entityCaptor = + ArgumentCaptor.forClass(StrategyGenerationEntity.class); + verify(repository).save(entityCaptor.capture()); + assertThat(entityCaptor.getValue().getKnowledgeBaseIds()).isEmpty(); + } + + @Test + @DisplayName("skipRetrieval=true 时普通用户也不做 KB 可见性校验") + void skipsVisibilityCheckWhenSkippingRetrieval() { + when(currentUserService.get()).thenReturn(new UserPrincipal(2L, "bob", UserRole.USER)); + stubChatAndSave(); + + service.generate(skipRequest()); + + verify(knowledgeBaseRepository, org.mockito.Mockito.never()) + .isVisibleToUser(any(), any()); + } + } + + @Nested + @DisplayName("显式用户入口") + class GenerateForUser { + + @Test + @DisplayName("generateForUser 使用传入用户生成,不依赖 SecurityContext") + void generatesWithProvidedUser() { + stubChatAndSave(); + + GenerateStrategyResponse response = service.generateForUser( + request(), new UserPrincipal(1L, "admin", UserRole.ADMIN)); + + assertThat(response).isNotNull(); + verify(currentUserService, org.mockito.Mockito.never()).get(); } } diff --git a/app/src/test/resources/application-test.yml b/app/src/test/resources/application-test.yml index 82ce767..52f0952 100644 --- a/app/src/test/resources/application-test.yml +++ b/app/src/test/resources/application-test.yml @@ -43,6 +43,8 @@ app: jwt: secret: test-jwt-secret-for-quantmore-tests-0123456789 expiration-days: 7 + eval: + enabled: false ai: default-provider: qwen security: