"""数据安全硬约束测试(对应 development-plan.md §7.2 / §10 纪律 1)。 这些测试是「禁止删除数据」「只写自有前缀」的**执行保证**, 不是靠自觉,而是靠断言。 """ from __future__ import annotations import re from pathlib import Path import pytest from hdiv.core.config import load_config from hdiv.core.errors import SafetyViolation from hdiv.core.paths import project_root from hdiv.data.db import StatementGuard, classify_statement, split_statements CFG = load_config("datasource") @pytest.fixture def guard() -> StatementGuard: return StatementGuard(CFG, allow_backfill=False) # --------------------------------------------------------------------------- # 语句分类 # --------------------------------------------------------------------------- @pytest.mark.parametrize( "sql,verb", [ ("SELECT 1", "select"), ("select * from stock", "select"), ("WITH x AS (SELECT 1) SELECT * FROM x", "with"), ("INSERT INTO hd_x VALUES (1)", "insert"), ("insert ignore into hd_x values (1)", "insert"), ("REPLACE INTO hd_x VALUES (1)", "replace"), ("UPDATE hd_x SET a=1", "update"), ("DELETE FROM hd_x", "delete"), ("CREATE TABLE hd_x (a INT)", "create"), ("ALTER TABLE hd_x ADD COLUMN b INT", "alter"), ("DROP TABLE hd_x", "drop"), ("TRUNCATE TABLE hd_x", "truncate"), ("-- 注释\nSELECT 1", "select"), ("/* 块注释 */ SELECT 1", "select"), ], ) def test_classify(sql: str, verb: str) -> None: assert classify_statement(sql)[0] == verb def test_split_statements() -> None: assert split_statements("SELECT 1; SELECT 2;") == ["SELECT 1", "SELECT 2"] assert split_statements("SELECT 1;;") == ["SELECT 1"] def test_classify_extracts_table() -> None: assert classify_statement("INSERT INTO `hd_dividend` (a) VALUES (1)")[1] == "hd_dividend" assert classify_statement("UPDATE stock SET a=1")[1] == "stock" assert classify_statement("DELETE FROM hd_x WHERE 1")[1] == "hd_x" # --------------------------------------------------------------------------- # 禁止删除:任何形式的 DELETE / DROP / TRUNCATE 都必须被拒 # --------------------------------------------------------------------------- @pytest.mark.parametrize( "sql", [ "DELETE FROM stock", "DELETE FROM stock WHERE symbol='000001.SZ'", "delete from hd_dividend where id=1", "TRUNCATE TABLE stock", "TRUNCATE stock", "DROP TABLE hd_x", "DROP TABLE IF EXISTS hd_x", "RENAME TABLE hd_x TO hd_y", "DELETE FROM stock_daily", "SELECT 1; DELETE FROM stock", "UPDATE hd_x SET a=1; DROP TABLE hd_y", ], ) def test_delete_like_always_blocked(guard: StatementGuard, sql: str) -> None: with pytest.raises(SafetyViolation): guard.check_script(sql) # --------------------------------------------------------------------------- # 只读白名单:qlib 既有表不可被写 # --------------------------------------------------------------------------- @pytest.mark.parametrize( "sql", [ "UPDATE stock SET name='x' WHERE symbol='000001.SZ'", "INSERT INTO stock_daily (symbol) VALUES ('000001.SZ')", "UPDATE daily_basic SET pe=1", "INSERT INTO financial_indicator (symbol) VALUES ('x')", "UPDATE trading_calendar SET is_open=0", "UPDATE stock_name_history SET name='x'", ], ) def test_readonly_tables_blocked(guard: StatementGuard, sql: str) -> None: with pytest.raises(SafetyViolation): guard.check_script(sql) @pytest.mark.parametrize( "sql", [ "SELECT * FROM stock", "SELECT * FROM stock_daily WHERE trade_date > '2020-01-01'", "WITH d AS (SELECT * FROM daily_basic) SELECT COUNT(*) FROM d", ], ) def test_readonly_tables_readable(guard: StatementGuard, sql: str) -> None: guard.check_script(sql) # 不应抛错 # --------------------------------------------------------------------------- # 自有前缀:只能写 hd_* 表 # --------------------------------------------------------------------------- @pytest.mark.parametrize( "sql", [ "INSERT INTO hd_dividend (symbol) VALUES ('x')", "UPDATE hd_strategy SET version='2'", "REPLACE INTO hd_sync_log (job) VALUES ('x')", ], ) def test_own_prefix_writes_allowed(guard: StatementGuard, sql: str) -> None: guard.check_script(sql) @pytest.mark.parametrize( "sql", [ "CREATE TABLE foo (a INT)", "CREATE TABLE tmp_data (a INT)", "ALTER TABLE stock_extra ADD COLUMN a INT", "INSERT INTO random_table (a) VALUES (1)", ], ) def test_foreign_table_writes_blocked(guard: StatementGuard, sql: str) -> None: with pytest.raises(SafetyViolation): guard.check_script(sql) def test_own_prefix_create_allowed(guard: StatementGuard) -> None: guard.check_script("CREATE TABLE IF NOT EXISTS hd_new_table (id INT)") # --------------------------------------------------------------------------- # 禁触表 # --------------------------------------------------------------------------- def test_alembic_version_protected(guard: StatementGuard) -> None: with pytest.raises(SafetyViolation): guard.check_script("UPDATE alembic_version SET version_num='x'") with pytest.raises(SafetyViolation): guard.check_script("ALTER TABLE alembic_version ADD COLUMN a INT") # --------------------------------------------------------------------------- # 回补通道:默认关闭,显式开启后只允许白名单表 # --------------------------------------------------------------------------- def test_backfill_requires_explicit_optin() -> None: g = StatementGuard(CFG, allow_backfill=False) with pytest.raises(SafetyViolation): g.check_script("INSERT INTO stock_daily (symbol) VALUES ('x')") def test_backfill_allows_only_declared_tables() -> None: g = StatementGuard(CFG, allow_backfill=True) # 已声明的表放行 g.check_script("INSERT IGNORE INTO stock_daily (symbol) VALUES ('x')") g.check_script("INSERT IGNORE INTO daily_basic (symbol) VALUES ('x')") # 未声明的表仍被拒 with pytest.raises(SafetyViolation): g.check_script("INSERT INTO financial_indicator (symbol) VALUES ('x')") def test_backfill_still_forbids_delete() -> None: g = StatementGuard(CFG, allow_backfill=True) with pytest.raises(SafetyViolation): g.check_script("DELETE FROM stock_daily") def test_index_weight_is_writable() -> None: """index_weight 是 qlib 的空表,经 allow_write_tables 允许纯新增填充。""" g = StatementGuard(CFG, allow_backfill=False) g.check_script("INSERT IGNORE INTO index_weight (index_code, trade_date, symbol) VALUES ('a','2020-01-01','b')") # --------------------------------------------------------------------------- # 源码级扫描:代码库里不得出现删除语句 # --------------------------------------------------------------------------- _FORBIDDEN_PATTERNS = [ (re.compile(r"\bDELETE\s+FROM\b", re.IGNORECASE), "DELETE FROM"), (re.compile(r"\bTRUNCATE\s+TABLE\b", re.IGNORECASE), "TRUNCATE TABLE"), (re.compile(r"\bDROP\s+TABLE\b", re.IGNORECASE), "DROP TABLE"), (re.compile(r"\.delete\s*\(", re.IGNORECASE), "session.delete("), (re.compile(r"\bDROP\s+INDEX\b(?!.*migrations)", re.IGNORECASE), "DROP INDEX"), ] # 允许出现这些词的文件(安全层自身、迁移定义、测试、文档) _ALLOWED_FILES = { "src/hdiv/data/db.py", "src/hdiv/data/schema.py", "src/hdiv/data/ddl.py", "src/hdiv/data/audit.py", "src/hdiv/data/sync/price.py", } def test_no_delete_statements_in_source() -> None: root = project_root() / "src" offenders: list[str] = [] for py in root.rglob("*.py"): rel = py.relative_to(project_root()).as_posix() if rel in _ALLOWED_FILES: continue text = py.read_text(encoding="utf-8") for pat, label in _FORBIDDEN_PATTERNS: for m in pat.finditer(text): line = text[: m.start()].count("\n") + 1 offenders.append(f"{rel}:{line} 含 {label}") assert not offenders, "源码中出现删除类语句:\n" + "\n".join(offenders) def test_no_qlib_code_import() -> None: """决策 D5:不 import qlib 项目代码,只共享数据库。""" root = project_root() / "src" bad: list[str] = [] pattern = re.compile(r"^\s*(?:from|import)\s+(app\.|qlib\.)", re.MULTILINE) for py in root.rglob("*.py"): text = py.read_text(encoding="utf-8") for m in pattern.finditer(text): line = text[: m.start()].count("\n") + 1 bad.append(f"{py.relative_to(project_root())}:{line}") assert not bad, f"禁止 import qlib 项目代码:{bad}" def test_guard_readonly_lists_are_configured() -> None: """白名单/前缀必须在配置里显式声明,不能靠默认值。""" db = CFG.database assert db.own_prefix == "hd_" assert "stock" in db.read_only_tables assert "stock_daily" in db.read_only_tables assert "alembic_version" in db.forbidden_tables assert db.forbid_delete is True assert Path(db.own_prefix) # 非空