import importlib.util
import unittest
from pathlib import Path


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


class D48TurtlePyramidMathTest(unittest.TestCase):
    def test_long_adds_at_half_and_one_atr(self):
        first = pyramid_math.pyramid_addition(
            initial_stake=120,
            initial_rate=100,
            initial_atr=4,
            current_rate=102,
            successful_entries=1,
            is_short=False,
            add_stake_fraction=0.25,
            spacing_atr=0.5,
            max_additions=2,
            min_stake=10,
            max_stake=100,
        )
        second = pyramid_math.pyramid_addition(
            initial_stake=120,
            initial_rate=100,
            initial_atr=4,
            current_rate=104,
            successful_entries=2,
            is_short=False,
            add_stake_fraction=0.25,
            spacing_atr=0.5,
            max_additions=2,
            min_stake=10,
            max_stake=100,
        )
        self.assertEqual((first.trigger_rate, first.stake), (102, 30))
        self.assertEqual((second.trigger_rate, second.stake), (104, 30))

    def test_short_threshold_is_symmetric(self):
        decision = pyramid_math.pyramid_addition(
            initial_stake=120,
            initial_rate=100,
            initial_atr=4,
            current_rate=98,
            successful_entries=1,
            is_short=True,
            add_stake_fraction=0.25,
            spacing_atr=0.5,
            max_additions=2,
            min_stake=None,
            max_stake=100,
        )
        self.assertEqual((decision.trigger_rate, decision.stake), (98, 30))

    def test_does_not_add_before_trigger_or_after_second_add(self):
        common = dict(
            initial_stake=120,
            initial_rate=100,
            initial_atr=4,
            is_short=False,
            add_stake_fraction=0.25,
            spacing_atr=0.5,
            max_additions=2,
            min_stake=None,
            max_stake=100,
        )
        self.assertIsNone(
            pyramid_math.pyramid_addition(
                current_rate=101.99, successful_entries=1, **common
            )
        )
        self.assertIsNone(
            pyramid_math.pyramid_addition(
                current_rate=110, successful_entries=3, **common
            )
        )

    def test_min_and_max_stake_are_enforced(self):
        capped = pyramid_math.pyramid_addition(
            initial_stake=120,
            initial_rate=100,
            initial_atr=4,
            current_rate=102,
            successful_entries=1,
            is_short=False,
            add_stake_fraction=0.25,
            spacing_atr=0.5,
            max_additions=2,
            min_stake=None,
            max_stake=20,
        )
        self.assertEqual(capped.stake, 20)
        self.assertIsNone(
            pyramid_math.pyramid_addition(
                initial_stake=120,
                initial_rate=100,
                initial_atr=4,
                current_rate=102,
                successful_entries=1,
                is_short=False,
                add_stake_fraction=0.25,
                spacing_atr=0.5,
                max_additions=2,
                min_stake=31,
                max_stake=100,
            )
        )

    def test_total_planned_risk_is_frozen_at_1_125_percent(self):
        risk = pyramid_math.planned_risk_fraction(
            initial_risk_fraction=0.0075,
            add_stake_fraction=0.25,
            max_additions=2,
        )
        self.assertAlmostEqual(risk, 0.01125)


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