import importlib.util
import unittest
from pathlib import Path


MODULE_PATH = (
    Path(__file__).resolve().parents[2]
    / "strategies"
    / "volatility_breakout"
    / "risk_cap_math.py"
)
SPEC = importlib.util.spec_from_file_location("risk_cap_math", MODULE_PATH)
risk_cap_math = importlib.util.module_from_spec(SPEC)
assert SPEC.loader is not None
SPEC.loader.exec_module(risk_cap_math)


class RiskCapMathTest(unittest.TestCase):
    def test_hard_stop_caps_low_atr_position(self):
        stake = risk_cap_math.risk_capped_stake(
            wallet=1000,
            risk_fraction=0.0075,
            atr_loss_fraction=0.03,
            hard_stop_fraction=0.05,
            wallet_cap_fraction=0.18,
            max_stake=500,
        )
        self.assertEqual(stake, 150)

    def test_atr_risk_still_controls_high_volatility_position(self):
        stake = risk_cap_math.risk_capped_stake(
            wallet=1000,
            risk_fraction=0.0075,
            atr_loss_fraction=0.06,
            hard_stop_fraction=0.05,
            wallet_cap_fraction=0.18,
            max_stake=500,
        )
        self.assertEqual(stake, 125)

    def test_exchange_max_stake_remains_a_hard_cap(self):
        stake = risk_cap_math.risk_capped_stake(
            wallet=1000,
            risk_fraction=0.0075,
            atr_loss_fraction=0.03,
            hard_stop_fraction=0.05,
            wallet_cap_fraction=0.18,
            max_stake=100,
        )
        self.assertEqual(stake, 100)

    def test_invalid_inputs_fail_closed(self):
        with self.assertRaises(ValueError):
            risk_cap_math.risk_capped_stake(
                wallet=1000,
                risk_fraction=0.0075,
                atr_loss_fraction=0,
                hard_stop_fraction=0.05,
                wallet_cap_fraction=0.18,
                max_stake=500,
            )


if __name__ == "__main__":
    unittest.main()
