diff --git a/backend/tests/test_backtest_warmup.py b/backend/tests/test_backtest_warmup.py index 3f00134..82bd37b 100644 --- a/backend/tests/test_backtest_warmup.py +++ b/backend/tests/test_backtest_warmup.py @@ -49,9 +49,17 @@ def test_load_panel_warms_up_indicators(monkeypatch) -> None: # warmup 行不进入结果面板 (pandas datetime64 与 date 直接比较会类型不符) assert str(panel["date"].min())[:10] == start.isoformat() assert str(panel["date"].max())[:10] == end.isoformat() - # 区间首日的指标已有历史窗口可用, 不再是 NaN + + # 数值复现: 区间头部的指标必须与「全历史计算后裁剪」的基准一致。 + # 只断言非 NaN 区分不了新旧代码 —— ewm 类指标 (RSI/MACD) 从首个值 + # 播种, 无 warmup 时首行也有值, 只是数值失真 (#201 的实际危害) + from app.indicators.pipeline import compute_all + reference = compute_all(df).filter(pl.col("date") >= start) + ref_rsi = float(reference["rsi_14"][0]) first = panel.iloc[0] - assert first["rsi_14"] == first["rsi_14"] # NaN != NaN + assert abs(first["rsi_14"] - ref_rsi) < 1.0, ( + f"rsi_14 without warmup: {first['rsi_14']} vs reference {ref_rsi}" + ) def test_load_panel_insufficient_history_degrades_gracefully(monkeypatch) -> None: diff --git a/backend/tests/test_screener_external_access.py b/backend/tests/test_screener_external_access.py index 6f4c2d8..6640e1d 100644 --- a/backend/tests/test_screener_external_access.py +++ b/backend/tests/test_screener_external_access.py @@ -37,7 +37,7 @@ def test_normal_condition_still_works() -> None: def test_injected_read_parquet_is_rejected() -> None: # 注入试图读任意文件: external_access 关闭后 DuckDB 直接报错, - # run_custom 的 except 分支吞错返回空结果, 而非泄漏文件内容 + # run() 的 except 分支吞错返回空结果, 而非泄漏文件内容 svc = _service_with_panel(_panel()) result = svc.run( date(2026, 9, 2), @@ -47,6 +47,25 @@ def test_injected_read_parquet_is_rejected() -> None: assert result.rows == [] +def test_injected_read_of_valid_parquet_leaks_nothing(tmp_path) -> None: + # 用合法 parquet 做注入目标才能区分新旧代码: /etc/passwd 不是 parquet, + # 未加固的连接读它也会报错; 合法同构文件在未加固连接下会真的混入结果 + victim = tmp_path / "victim.parquet" + _panel().write_parquet(victim) + svc = _service_with_panel(pl.DataFrame( + {"symbol": ["999999.SZ"], "close": [99.0], "turnover_rate": [9.0]} + )) + result = svc.run( + date(2026, 9, 2), + [f"1=1) UNION SELECT * FROM read_parquet('{victim}') --"], + limit=10, + ) + # 加固后注入被拒 → 整条查询 fail-closed 返回空 (或至多剩原面板行); + # 未加固时 victim 的 2 行会混入结果 + assert all(r["symbol"] == "999999.SZ" for r in result.rows) + assert len(result.rows) <= 1 + + def test_injected_copy_write_is_rejected(tmp_path) -> None: target = tmp_path / "pwned.csv" svc = _service_with_panel(_panel())