mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-07-27 23:47:46 +00:00
test: add direct unit tests for _apply_riskmgmt_mask, safe_resolve_path, import_class, and _add_column_if_not_exists
This commit is contained in:
@@ -16,7 +16,9 @@ import pytest
|
||||
|
||||
from rdagent.components.backtesting.vbt_backtest import (
|
||||
OOS_START_DEFAULT,
|
||||
_apply_ftmo_mask,
|
||||
backtest_signal_ftmo,
|
||||
FTMO_INITIAL_CAPITAL,
|
||||
FTMO_MAX_DAILY_LOSS,
|
||||
FTMO_MAX_TOTAL_LOSS,
|
||||
monte_carlo_trade_pvalue,
|
||||
@@ -238,3 +240,141 @@ def test_wf_consistency_range(close_6yr):
|
||||
c = r.get("wf_oos_consistency")
|
||||
if c is not None:
|
||||
assert 0.0 <= c <= 1.0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Direct _apply_ftmo_mask unit tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestApplyFtmoMask:
|
||||
"""Direct unit tests for _apply_ftmo_mask — the core FTMO daily/total loss engine."""
|
||||
|
||||
@pytest.fixture
|
||||
def flat_close(self) -> pd.Series:
|
||||
n = 3000
|
||||
idx = pd.date_range("2024-01-01", periods=n, freq="1min")
|
||||
return pd.Series(1.10, index=idx)
|
||||
|
||||
def test_returns_compliance_dict(self, flat_close):
|
||||
signal = _random_signal(flat_close.index)
|
||||
masked, info = _apply_ftmo_mask(signal, flat_close, leverage=1.0, txn_cost_bps=2.14)
|
||||
assert "ftmo_daily_breaches" in info
|
||||
assert "ftmo_total_breached" in info
|
||||
assert "ftmo_total_breach_ts" in info
|
||||
assert "ftmo_compliant" in info
|
||||
|
||||
def test_flat_market_zero_signal_fully_compliant(self, flat_close):
|
||||
"""No trades → always compliant."""
|
||||
signal = pd.Series(0.0, index=flat_close.index)
|
||||
masked, info = _apply_ftmo_mask(signal, flat_close, leverage=1.0, txn_cost_bps=2.14)
|
||||
assert info["ftmo_daily_breaches"] == 0
|
||||
assert info["ftmo_total_breached"] is False
|
||||
assert info["ftmo_compliant"] is True
|
||||
# All signals should remain zero
|
||||
assert (masked == 0).all()
|
||||
|
||||
def test_daily_loss_breach_zeroes_rest_of_day(self):
|
||||
"""When daily loss exceeds 5%, rest of that day's signals are zeroed."""
|
||||
n = 3000
|
||||
idx = pd.date_range("2024-01-01", periods=n, freq="1min")
|
||||
# Price drops sharply in first few bars to trigger daily loss
|
||||
price = pd.Series(1.10, index=idx, dtype=float)
|
||||
price.iloc[3:20] = 0.00 # crash from 1.10 to 0.00 → massive loss
|
||||
signal = pd.Series(1.0, index=idx) # always long at 30x leverage
|
||||
|
||||
masked, info = _apply_ftmo_mask(signal, price, leverage=30.0, txn_cost_bps=0)
|
||||
assert info["ftmo_daily_breaches"] > 0
|
||||
# After breach, signals on same day must be zeroed
|
||||
breach_day = idx[0].date()
|
||||
same_day_late = (idx[-1] if idx[-1].date() == breach_day else idx[20])
|
||||
if same_day_late.date() == breach_day:
|
||||
assert masked.loc[same_day_late] == 0
|
||||
|
||||
def test_total_loss_breach_zeroes_all_remaining(self):
|
||||
"""When total loss exceeds 10%, ALL subsequent signals are zeroed."""
|
||||
n = 5000
|
||||
idx = pd.date_range("2024-01-01", periods=n, freq="1min")
|
||||
# Price crashes → max position → total loss limit breached
|
||||
price = pd.Series(1.10, index=idx, dtype=float)
|
||||
price.iloc[5:50] = 0.50 # >10% drop with 30x leverage
|
||||
signal = pd.Series(1.0, index=idx)
|
||||
|
||||
masked, info = _apply_ftmo_mask(signal, price, leverage=30.0, txn_cost_bps=0)
|
||||
assert info["ftmo_total_breached"] is True
|
||||
assert info["ftmo_total_breach_ts"] is not None
|
||||
# After breach, ALL later signals must be zero
|
||||
assert (masked.iloc[100:] == 0).all()
|
||||
|
||||
def test_total_breach_respected_across_days(self):
|
||||
"""Total breach persists across day boundaries — no new trades after breach."""
|
||||
n = 5000
|
||||
idx = pd.date_range("2024-01-01", periods=n, freq="1min")
|
||||
price = pd.Series(1.10, index=idx, dtype=float)
|
||||
price.iloc[5:50] = 0.50
|
||||
signal = pd.Series(1.0, index=idx)
|
||||
|
||||
masked, info = _apply_ftmo_mask(signal, price, leverage=30.0, txn_cost_bps=0)
|
||||
# All signals after breach index must be zero
|
||||
breach_ts = pd.Timestamp(info["ftmo_total_breach_ts"])
|
||||
assert (masked.loc[masked.index > breach_ts] == 0).all()
|
||||
|
||||
def test_daily_loss_resets_on_new_day(self):
|
||||
"""Daily loss limit resets at day boundary — new day starts fresh (unless total breached)."""
|
||||
n = 5000
|
||||
idx = pd.date_range("2024-01-01", periods=n, freq="1min")
|
||||
price = pd.Series(1.10, index=idx, dtype=float)
|
||||
# Trigger daily breach on day 1 by dropping 1%
|
||||
price.iloc[5:20] = 1.09 # ~1% drop with 30x → ~30% loss
|
||||
signal = pd.Series(1.0, index=idx)
|
||||
|
||||
masked, info = _apply_ftmo_mask(signal, price, leverage=30.0, txn_cost_bps=0)
|
||||
assert info["ftmo_daily_breaches"] >= 1
|
||||
# Day 2 signals should be active again if not total-breached
|
||||
day2_mask = idx.date > idx[0].date()
|
||||
if day2_mask.any() and not info["ftmo_total_breached"]:
|
||||
day2 = idx[day2_mask][0]
|
||||
assert masked.loc[day2] != 0
|
||||
|
||||
def test_compliant_flag_false_after_daily_breach(self):
|
||||
"""Even one daily breach makes ftmo_compliant=False."""
|
||||
n = 3000
|
||||
idx = pd.date_range("2024-01-01", periods=n, freq="1min")
|
||||
price = pd.Series(1.10, index=idx, dtype=float)
|
||||
price.iloc[3:20] = 0.00
|
||||
signal = pd.Series(1.0, index=idx)
|
||||
|
||||
masked, info = _apply_ftmo_mask(signal, price, leverage=30.0, txn_cost_bps=0)
|
||||
assert info["ftmo_compliant"] is False
|
||||
|
||||
def test_compliant_flag_false_after_total_breach(self):
|
||||
"""Total breach makes ftmo_compliant=False."""
|
||||
n = 5000
|
||||
idx = pd.date_range("2024-01-01", periods=n, freq="1min")
|
||||
price = pd.Series(1.10, index=idx, dtype=float)
|
||||
price.iloc[5:50] = 0.50
|
||||
signal = pd.Series(1.0, index=idx)
|
||||
|
||||
masked, info = _apply_ftmo_mask(signal, price, leverage=30.0, txn_cost_bps=0)
|
||||
assert info["ftmo_compliant"] is False
|
||||
|
||||
def test_transaction_costs_reduce_equity(self):
|
||||
"""Transaction costs should reduce equity — compliant scenario with fees."""
|
||||
n = 1000
|
||||
idx = pd.date_range("2024-01-01", periods=n, freq="1min")
|
||||
price = pd.Series(1.10, index=idx, dtype=float)
|
||||
# Alternating signal → lots of position changes → high costs
|
||||
signal = pd.Series([1.0 if i % 2 == 0 else -1.0 for i in range(n)], index=idx)
|
||||
|
||||
masked, info = _apply_ftmo_mask(signal, price, leverage=1.0, txn_cost_bps=10.0)
|
||||
# With high costs and flat market, equity should drop
|
||||
assert "ftmo_daily_breaches" in info
|
||||
|
||||
def test_output_mask_has_same_index(self):
|
||||
n = 2000
|
||||
idx = pd.date_range("2024-01-01", periods=n, freq="1min")
|
||||
price = pd.Series(1.10, index=idx)
|
||||
signal = _random_signal(idx, seed=1)
|
||||
|
||||
masked, info = _apply_ftmo_mask(signal, price, leverage=1.0, txn_cost_bps=2.14)
|
||||
assert len(masked) == len(signal)
|
||||
assert masked.index.equals(signal.index)
|
||||
|
||||
@@ -399,3 +399,91 @@ class TestDatabaseIntegrity:
|
||||
|
||||
# Import am Anfang der Datei für die Tests
|
||||
from rdagent.components.backtesting.results_db import ResultsDatabase
|
||||
|
||||
|
||||
class TestAddColumnIfNotExists:
|
||||
"""Direct tests for _add_column_if_not_exists migration helper."""
|
||||
|
||||
def test_add_new_column_succeeds(self):
|
||||
"""Adding a new column to an existing table should work."""
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
import os
|
||||
db_path = os.path.join(tmpdir, "test.db")
|
||||
db = ResultsDatabase(db_path=db_path)
|
||||
try:
|
||||
db._add_column_if_not_exists("backtest_runs", "test_new_col", "REAL")
|
||||
c = db.conn.cursor()
|
||||
c.execute("PRAGMA table_info(backtest_runs)")
|
||||
cols = [row[1] for row in c.fetchall()]
|
||||
assert "test_new_col" in cols
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
def test_existing_column_noop(self):
|
||||
"""Adding an already existing column should succeed (no-op)."""
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
import os
|
||||
db_path = os.path.join(tmpdir, "test.db")
|
||||
db = ResultsDatabase(db_path=db_path)
|
||||
try:
|
||||
# First call adds, second call should be no-op
|
||||
db._add_column_if_not_exists("backtest_runs", "ic", "REAL")
|
||||
db._add_column_if_not_exists("backtest_runs", "ic", "REAL")
|
||||
c = db.conn.cursor()
|
||||
c.execute("PRAGMA table_info(backtest_runs)")
|
||||
cols = [row[1] for row in c.fetchall()]
|
||||
assert cols.count("ic") == 1 # should exist exactly once
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
def test_invalid_table_raises(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
import os
|
||||
db_path = os.path.join(tmpdir, "test.db")
|
||||
db = ResultsDatabase(db_path=db_path)
|
||||
try:
|
||||
with pytest.raises(ValueError, match="Unknown table"):
|
||||
db._add_column_if_not_exists("nonexistent_table", "col", "REAL")
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
def test_invalid_column_name_raises(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
import os
|
||||
db_path = os.path.join(tmpdir, "test.db")
|
||||
db = ResultsDatabase(db_path=db_path)
|
||||
try:
|
||||
with pytest.raises(ValueError, match="Invalid column name"):
|
||||
db._add_column_if_not_exists("backtest_runs", "bad;column", "REAL")
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
def test_invalid_column_type_raises(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
import os
|
||||
db_path = os.path.join(tmpdir, "test.db")
|
||||
db = ResultsDatabase(db_path=db_path)
|
||||
try:
|
||||
with pytest.raises(ValueError, match="Invalid column type"):
|
||||
db._add_column_if_not_exists("backtest_runs", "col", "INVALID_TYPE")
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
def test_all_allowed_types_work(self):
|
||||
"""REAL, TEXT, INTEGER, BLOB should all be valid types."""
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
import os
|
||||
db_path = os.path.join(tmpdir, "test.db")
|
||||
db = ResultsDatabase(db_path=db_path)
|
||||
try:
|
||||
for col_type in ("REAL", "TEXT", "INTEGER", "BLOB"):
|
||||
db._add_column_if_not_exists(
|
||||
"backtest_runs", f"test_{col_type.lower()}", col_type,
|
||||
)
|
||||
c = db.conn.cursor()
|
||||
c.execute("PRAGMA table_info(backtest_runs)")
|
||||
cols = {row[1] for row in c.fetchall()}
|
||||
for col_type in ("REAL", "TEXT", "INTEGER", "BLOB"):
|
||||
assert f"test_{col_type.lower()}" in cols
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
+74
-1
@@ -1,8 +1,10 @@
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from rdagent.core.utils import SingletonBaseClass
|
||||
from rdagent.core.utils import SingletonBaseClass, import_class, safe_resolve_path
|
||||
|
||||
|
||||
class A(SingletonBaseClass):
|
||||
@@ -70,5 +72,76 @@ class MiscTest(unittest.TestCase):
|
||||
# print(a1.kwargs) # a1 will be changed.
|
||||
|
||||
|
||||
class TestSafeResolvePath:
|
||||
"""Tests for safe_resolve_path — path traversal prevention."""
|
||||
|
||||
def test_inside_root_returns_absolute(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
root = Path(tmpdir)
|
||||
result = safe_resolve_path(root / "subdir" / "file.txt", safe_root=root)
|
||||
assert result.is_absolute()
|
||||
assert str(result).startswith(str(root.resolve()))
|
||||
|
||||
def test_no_safe_root_just_resolves(self):
|
||||
result = safe_resolve_path(Path("/tmp/nonexistent_test"), safe_root=None)
|
||||
assert result.is_absolute()
|
||||
|
||||
def test_path_traversal_raises(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
root = Path(tmpdir)
|
||||
with pytest.raises(ValueError, match="outside allowed root"):
|
||||
safe_resolve_path(root / ".." / "etc" / "passwd", safe_root=root)
|
||||
|
||||
def test_symlink_outside_root_raises(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
root = Path(tmpdir)
|
||||
inside = root / "inside"
|
||||
inside.mkdir()
|
||||
link = inside / "escape"
|
||||
link.symlink_to("/etc/passwd")
|
||||
with pytest.raises(ValueError, match="outside allowed root"):
|
||||
safe_resolve_path(link, safe_root=root)
|
||||
|
||||
def test_root_itself_is_valid(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
root = Path(tmpdir)
|
||||
result = safe_resolve_path(root, safe_root=root)
|
||||
assert result == root.resolve()
|
||||
|
||||
def test_expanduser_resolves_home(self):
|
||||
result = safe_resolve_path(Path("~/nonexistent_test"), safe_root=None)
|
||||
assert str(result).startswith(str(Path.home()))
|
||||
|
||||
|
||||
class TestImportClass:
|
||||
"""Tests for import_class — dynamic class loading."""
|
||||
|
||||
def test_valid_class_import(self):
|
||||
cls = import_class("pathlib.Path")
|
||||
assert cls is Path
|
||||
|
||||
def test_builtin_class_import(self):
|
||||
cls = import_class("collections.OrderedDict")
|
||||
from collections import OrderedDict
|
||||
assert cls is OrderedDict
|
||||
|
||||
def test_invalid_module_raises_import_error(self):
|
||||
with pytest.raises(ImportError, match="Module not found"):
|
||||
import_class("nonexistent.module.ClassName")
|
||||
|
||||
def test_missing_class_raises_import_error(self):
|
||||
with pytest.raises(ImportError, match="Class not found"):
|
||||
import_class("pathlib.NonExistentClass")
|
||||
|
||||
def test_invalid_format_raises_import_error(self):
|
||||
with pytest.raises(ImportError, match="Invalid class path"):
|
||||
import_class("no_dots_at_all")
|
||||
|
||||
def test_pandas_class_import(self):
|
||||
cls = import_class("pandas.DataFrame")
|
||||
import pandas as pd
|
||||
assert cls is pd.DataFrame
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user