200 lines
6.4 KiB
Python
200 lines
6.4 KiB
Python
"""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<Long> ids) {
|
|
for (Long id : ids) {
|
|
FileTaskEntity task = fileTaskMapper.selectById(id);
|
|
}
|
|
}
|
|
}
|
|
"""
|
|
|
|
FIXTURE_QUERY_RESULT_LOOP = """package com.example;
|
|
public class DemoService {
|
|
public List<Vo> listVos() {
|
|
List<FileTaskEntity> tasks = fileTaskMapper.selectList(new LambdaQueryWrapper<>());
|
|
List<Vo> 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<Long> ids) {
|
|
ids.forEach(id -> fileTaskMapper.selectOne(
|
|
new LambdaQueryWrapper<FileTaskEntity>().eq(FileTaskEntity::getId, id)));
|
|
}
|
|
}
|
|
"""
|
|
|
|
FIXTURE_BATCH_SAFE = """package com.example;
|
|
public class DemoService {
|
|
public List<FileTaskEntity> load(List<Long> ids) {
|
|
List<FileTaskEntity> 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<Integer> nums) {
|
|
int s = 0;
|
|
for (Integer n : nums) {
|
|
s += n;
|
|
}
|
|
return s;
|
|
}
|
|
}
|
|
"""
|
|
|
|
FIXTURE_INDEX_LOOP = """package com.example;
|
|
public class DemoService {
|
|
public List<Vo> listVos(List<Long> ids) {
|
|
List<Vo> 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()
|