"""task-121 N+1 扫描审计测试。 对应 plan 06 任务 121 的 8 条用例 + 2 条补强: 1. test_scan_select_in_loop 识别循环内单查 2. test_scan_output_listed 清单产出(位置/对象/影响字段齐全) 3. test_scan_priority 按影响排序 4. test_scan_module_coverage 覆盖主要模块(真实代码扫描) 5. test_scan_false_positive_checked 排除误报(已批量/循环外) 6. test_scan_repeatable 可重复(同输入同输出) 7. test_scan_documented 审计文档产出 8. test_scan_no_edit 本任务不改生产代码 9. test_scan_lambda_brace_less 无花括号 lambda 内单查识别 10. test_scan_iterating_query_result 遍历查询结果集判 HIGH """ from __future__ import annotations import json import subprocess import unittest from pathlib import Path from n1_scan import ( HIGH, LOW, MEDIUM, MODULES_ROOT, scan_modules, scan_text, ) SCRIPTS_DIR = Path(__file__).resolve().parent REPO_ROOT = SCRIPTS_DIR.parent.parent DOC_PATH = REPO_ROOT / "backend-java" / "docs" / "n1-scan-audit.md" FIXTURE_LOOP = """package com.example; public class DemoService { public void batchByIds(List ids) { for (Long id : ids) { FileTaskEntity task = fileTaskMapper.selectById(id); } } } """ FIXTURE_QUERY_RESULT_LOOP = """package com.example; public class DemoService { public List listVos() { List tasks = fileTaskMapper.selectList(new LambdaQueryWrapper<>()); List vos = new ArrayList<>(); for (FileTaskEntity t : tasks) { FileResultEntity r = fileResultMapper.selectById(t.getResultId()); vos.add(toVo(r)); } return vos; } } """ FIXTURE_LAMBDA_BRACELESS = """package com.example; public class DemoService { public void touch(List ids) { ids.forEach(id -> fileTaskMapper.selectOne( new LambdaQueryWrapper().eq(FileTaskEntity::getId, id))); } } """ FIXTURE_BATCH_SAFE = """package com.example; public class DemoService { public List load(List ids) { List tasks = new ArrayList<>(); for (Long id : ids) { tasks.addAll(fileTaskMapper.selectBatchIds(Collections.singleton(id))); } for (Long id : ids) { tasks.addAll(fileTaskMapper.selectByIds(Collections.singleton(id))); } return tasks; } } """ FIXTURE_NO_LOOP = """package com.example; public class DemoService { public FileTaskEntity one(Long id) { return fileTaskMapper.selectById(id); } } """ FIXTURE_LOOP_NO_QUERY = """package com.example; public class DemoService { public int sum(List nums) { int s = 0; for (Integer n : nums) { s += n; } return s; } } """ FIXTURE_INDEX_LOOP = """package com.example; public class DemoService { public List listVos(List ids) { List vos = new ArrayList<>(); for (int i = 0; i < ids.size(); i++) { FileTaskEntity t = fileTaskMapper.selectById(ids.get(i)); vos.add(toVo(t)); } return vos; } } """ FIXTURE_MIXED = ( FIXTURE_QUERY_RESULT_LOOP + "\n" + FIXTURE_LAMBDA_BRACELESS + "\n" + FIXTURE_INDEX_LOOP ) class ScanDetectTest(unittest.TestCase): def test_scan_select_in_loop(self): findings = scan_text(FIXTURE_LOOP, "DemoService.java") self.assertEqual(len(findings), 1) self.assertIn("selectById", findings[0]["statement"]) self.assertEqual(findings[0]["loop_type"], "for") self.assertEqual(findings[0]["line"], 5) def test_scan_output_listed(self): for finding in scan_text(FIXTURE_MIXED, "DemoService.java"): for key in ("file", "line", "statement", "loop_type", "impact"): self.assertIn(key, finding) self.assertTrue(finding[key] is not None and finding[key] != "") def test_scan_priority(self): findings = scan_text(FIXTURE_MIXED, "DemoService.java") order = {HIGH: 0, MEDIUM: 1, LOW: 2} impacts = [order[f["impact"]] for f in findings] self.assertEqual(impacts, sorted(impacts)) def test_scan_false_positive_checked(self): self.assertEqual(len(scan_text(FIXTURE_BATCH_SAFE, "Batch.java")), 0) self.assertEqual(len(scan_text(FIXTURE_NO_LOOP, "NoLoop.java")), 0) self.assertEqual(len(scan_text(FIXTURE_LOOP_NO_QUERY, "NoQuery.java")), 0) def test_scan_repeatable(self): first = scan_text(FIXTURE_MIXED, "DemoService.java") second = scan_text(FIXTURE_MIXED, "DemoService.java") self.assertEqual( json.dumps(first, sort_keys=True), json.dumps(second, sort_keys=True) ) def test_scan_lambda_brace_less(self): findings = scan_text(FIXTURE_LAMBDA_BRACELESS, "DemoService.java") self.assertEqual(len(findings), 1) self.assertEqual(findings[0]["loop_type"], "forEach") self.assertEqual(findings[0]["impact"], MEDIUM) def test_scan_iterating_query_result(self): findings = scan_text(FIXTURE_QUERY_RESULT_LOOP, "DemoService.java") self.assertEqual(len(findings), 1) self.assertEqual(findings[0]["impact"], HIGH) class ScanCoverageTest(unittest.TestCase): def test_scan_module_coverage(self): report = scan_modules(MODULES_ROOT) self.assertGreaterEqual(report.scanned_files, 600) self.assertGreater(len(report.findings), 0) modules = {f["module"] for f in report.findings} self.assertGreaterEqual(len(modules), 3) order = {HIGH: 0, MEDIUM: 1, LOW: 2} impacts = [order[f["impact"]] for f in report.findings] self.assertEqual(impacts, sorted(impacts)) def test_scan_documented(self): self.assertTrue(DOC_PATH.is_file(), f"audit doc missing: {DOC_PATH}") text = DOC_PATH.read_text(encoding="utf-8") self.assertIn("N+1", text) self.assertIn("影响", text) def test_scan_no_edit(self): changed = subprocess.run( ["git", "diff", "--name-only", "HEAD", "--", "backend-java/src"], cwd=REPO_ROOT, capture_output=True, text=True, check=True, ).stdout.strip() self.assertEqual(changed, "", f"生产代码被改动: {changed}") if __name__ == "__main__": unittest.main()