"""
Known-case tests for the point-in-time replay engine (StockLab final engineering pass, Part C).

Pure, dependency-free — genuinely executed. Dual-mode: pytest, or
`PYTHONPATH=. python3 tests/test_replay.py`.

**What these tests establish and what they do not.** They establish that the replay LOOP is
correct: that it realises the right period returns from a known price path, that it refuses
look-ahead data, that it flags survivorship/adjustment/cost risks, and that turnover and costs are
charged correctly. They establish nothing whatsoever about StockLab's strategy — the price paths
below are hand-built arithmetic, not market data, and no result from them is a backtest result.
"""
from __future__ import annotations

from datetime import date

from app.engines.backtesting.performance import compute_performance_report
from app.engines.backtesting.replay import (
    LookAheadBiasError,
    PriceObservation,
    run_replay,
)

D1, D2, D3 = date(2020, 1, 1), date(2020, 4, 1), date(2020, 7, 1)


class FakeSource:
    """An in-memory SnapshotSource. `prices[date][security_id] = price`."""

    def __init__(self, prices, universes=None, known_as_of=None, adjusted=True):
        self.prices = prices
        self.universes = universes
        self.known_as_of = known_as_of or {}
        self.adjusted = adjusted

    def universe_at(self, as_of):
        if self.universes is not None:
            return self.universes[as_of]
        return sorted(self.prices.get(as_of, {}))

    def observe(self, security_id, as_of):
        price = self.prices.get(as_of, {}).get(security_id)
        if price is None:
            return None
        return PriceObservation(
            security_id=security_id, price=price,
            known_as_of=self.known_as_of.get((security_id, as_of), as_of),
            is_split_adjusted=self.adjusted,
        )


def _hold_all(_obs):
    return 1.0


# --- the core loop ---


def test_single_security_period_return_is_the_price_return():
    src = FakeSource({D1: {"A": 100.0}, D2: {"A": 110.0}})
    r = run_replay(src, _hold_all, [D1, D2])
    assert r.periods == 1
    assert abs(r.period_returns[0] - 0.10) < 1e-12


def test_two_securities_are_equally_weighted_by_default():
    """A +20% and a -10% at equal weight is +5%."""
    src = FakeSource({D1: {"A": 100.0, "B": 100.0}, D2: {"A": 120.0, "B": 90.0}})
    r = run_replay(src, _hold_all, [D1, D2])
    assert abs(r.period_returns[0] - 0.05) < 1e-12


def test_weights_are_normalised_not_taken_literally():
    """decide() returning 3.0 and 1.0 means 75%/25%, not 300%/100%."""
    def decide(obs):
        return 3.0 if obs.security_id == "A" else 1.0

    src = FakeSource({D1: {"A": 100.0, "B": 100.0}, D2: {"A": 200.0, "B": 100.0}})
    r = run_replay(src, decide, [D1, D2])
    assert abs(r.period_returns[0] - 0.75) < 1e-12  # 0.75 * 100% + 0.25 * 0%


def test_zero_weight_security_is_not_held():
    def decide(obs):
        return 0.0 if obs.security_id == "B" else 1.0

    src = FakeSource({D1: {"A": 100.0, "B": 100.0}, D2: {"A": 110.0, "B": 500.0}})
    r = run_replay(src, decide, [D1, D2])
    assert abs(r.period_returns[0] - 0.10) < 1e-12  # B's +400% is not in the book
    assert set(r.rebalances[0].holdings) == {"A"}


def test_an_all_zero_book_holds_cash_not_an_equal_weight_portfolio():
    src = FakeSource({D1: {"A": 100.0}, D2: {"A": 200.0}})
    r = run_replay(src, lambda _o: 0.0, [D1, D2])
    assert r.rebalances[0].holdings == {}
    assert r.period_returns[0] == 0.0


def test_returns_compound_across_multiple_periods():
    src = FakeSource({D1: {"A": 100.0}, D2: {"A": 110.0}, D3: {"A": 121.0}})
    r = run_replay(src, _hold_all, [D1, D2, D3])
    assert r.periods == 2
    assert all(abs(x - 0.10) < 1e-12 for x in r.period_returns)
    report = compute_performance_report(r.period_returns, periods_per_year=4.0)
    assert abs(report.total_return - 0.21) < 1e-12


def test_rebalance_dates_are_sorted_and_deduplicated():
    src = FakeSource({D1: {"A": 100.0}, D2: {"A": 110.0}})
    r = run_replay(src, _hold_all, [D2, D1, D2])
    assert [rec.as_of for rec in r.rebalances] == [D1, D2]


def test_fewer_than_two_dates_produces_no_periods_and_says_why():
    src = FakeSource({D1: {"A": 100.0}})
    r = run_replay(src, _hold_all, [D1])
    assert r.periods == 0
    assert any("insufficient_rebalance_dates" in w for w in r.warnings)


# --- look-ahead bias: the one thing the engine enforces mechanically ---


def test_observation_dated_after_the_rebalance_raises():
    src = FakeSource(
        {D1: {"A": 100.0}, D2: {"A": 110.0}},
        known_as_of={("A", D1): date(2020, 2, 1)},  # a month AFTER the D1 decision
    )
    try:
        run_replay(src, _hold_all, [D1, D2])
    except LookAheadBiasError as exc:
        assert "known_as_of" in str(exc)
        return
    raise AssertionError("look-ahead data was accepted")


def test_observation_dated_exactly_on_the_rebalance_date_is_allowed():
    src = FakeSource({D1: {"A": 100.0}, D2: {"A": 110.0}},
                     known_as_of={("A", D1): D1})
    assert run_replay(src, _hold_all, [D1, D2]).periods == 1


def test_observation_dated_before_the_rebalance_date_is_allowed():
    src = FakeSource({D1: {"A": 100.0}, D2: {"A": 110.0}},
                     known_as_of={("A", D1): date(2019, 11, 15)})
    assert run_replay(src, _hold_all, [D1, D2]).periods == 1


# --- survivorship, price adjustment, costs ---


def test_a_constant_universe_raises_the_survivorship_suspicion():
    src = FakeSource({D1: {"A": 100.0}, D2: {"A": 110.0}})
    r = run_replay(src, _hold_all, [D1, D2])
    assert r.survivorship_bias_suspected
    assert any("universe_constant" in w for w in r.warnings)
    assert not r.is_quotable


def test_a_varying_universe_does_not_raise_the_suspicion():
    src = FakeSource(
        {D1: {"A": 100.0, "B": 50.0}, D2: {"A": 110.0}},
        universes={D1: ["A", "B"], D2: ["A"]},
    )
    r = run_replay(src, _hold_all, [D1, D2])
    assert not r.survivorship_bias_suspected


def test_unadjusted_prices_are_flagged():
    src = FakeSource({D1: {"A": 100.0}, D2: {"A": 110.0}}, adjusted=False)
    r = run_replay(src, _hold_all, [D1, D2])
    assert r.unadjusted_prices_used
    assert not r.is_quotable


def test_zero_costs_are_flagged_as_not_modelled():
    src = FakeSource({D1: {"A": 100.0}, D2: {"A": 110.0}})
    r = run_replay(src, _hold_all, [D1, D2])
    assert not r.costs_modelled
    assert any("cost_per_unit_turnover=0" in w for w in r.warnings)


def test_turnover_cost_is_deducted_from_the_period_return():
    """Costs are charged at the rebalance that CAUSES the turnover, against the return realised at
    that same date. Here the book is fully in A from D1, so the D1->D2 rebalance has zero turnover
    and zero cost; the D3 rebalance adds B (turnover 1.0) and pays for it out of the D2->D3
    period."""
    src = FakeSource(
        {D1: {"A": 100.0}, D2: {"A": 110.0}, D3: {"A": 121.0, "B": 50.0}},
        universes={D1: ["A"], D2: ["A"], D3: ["A", "B"]},
    )
    gross = run_replay(src, _hold_all, [D1, D2, D3])
    net = run_replay(src, _hold_all, [D1, D2, D3], cost_per_unit_turnover=0.01)
    assert net.costs_modelled
    # period 0 (D1->D2): turnover 0 at D2 -> unchanged
    assert abs(net.period_returns[0] - gross.period_returns[0]) < 1e-12
    assert abs(net.turnover_per_period[0]) < 1e-12
    # period 1 (D2->D3): turnover 1.0 at D3 (half the book rotates into B) -> costs 1.0 * 0.01
    assert abs(net.turnover_per_period[1] - 1.0) < 1e-12
    assert abs((gross.period_returns[1] - net.period_returns[1]) - 0.01) < 1e-12


def test_first_rebalance_turnover_is_recorded_even_though_it_realises_no_return():
    src = FakeSource({D1: {"A": 100.0}, D2: {"A": 110.0}})
    r = run_replay(src, _hold_all, [D1, D2])
    assert abs(r.rebalances[0].turnover - 1.0) < 1e-12
    assert r.rebalances[0].period_return is None


def test_switching_the_whole_book_costs_two_units_of_turnover():
    src = FakeSource(
        {D1: {"A": 100.0, "B": 100.0}, D2: {"A": 110.0, "B": 110.0}, D3: {"A": 120.0, "B": 120.0}},
        universes={D1: ["A"], D2: ["B"], D3: ["A", "B"]},
    )
    r = run_replay(src, _hold_all, [D1, D2, D3])
    assert abs(r.rebalances[1].turnover - 2.0) < 1e-12  # sell all of A, buy all of B


# --- delisting / missing price handling ---


def test_a_held_security_that_loses_its_price_is_held_flat_and_reported():
    """Dropping it instead would silently re-weight the portfolio onto its survivors -- exactly the
    bias this engine exists to avoid, reintroduced at the smallest scale."""
    src = FakeSource(
        {D1: {"A": 100.0, "B": 100.0}, D2: {"A": 120.0}},
        universes={D1: ["A", "B"], D2: ["A"]},
    )
    r = run_replay(src, _hold_all, [D1, D2])
    # A returned +20% at 50% weight; B contributes 0%, so the period is +10%, not +20%.
    assert abs(r.period_returns[0] - 0.10) < 1e-12
    assert any("no price for 1 held securities" in w for w in r.warnings)


def test_a_security_with_no_observation_is_never_bought():
    class NoneSource(FakeSource):
        def observe(self, security_id, as_of):
            if security_id == "B":
                return None
            return super().observe(security_id, as_of)

    src = NoneSource(
        {D1: {"A": 100.0, "B": 100.0}, D2: {"A": 110.0, "B": 100.0}},
        universes={D1: ["A", "B"], D2: ["A"]},
    )
    r = run_replay(src, _hold_all, [D1, D2])
    assert set(r.rebalances[0].holdings) == {"A"}
    assert r.rebalances[0].securities_considered == 2
    assert r.rebalances[0].securities_priced == 1


# --- the integrity gate ---


def test_is_quotable_requires_every_flag_to_be_clear():
    src = FakeSource(
        {D1: {"A": 100.0, "B": 100.0}, D2: {"A": 110.0}},
        universes={D1: ["A", "B"], D2: ["A"]},
    )
    r = run_replay(src, _hold_all, [D1, D2], cost_per_unit_turnover=0.001)
    assert not r.survivorship_bias_suspected
    assert not r.unadjusted_prices_used
    assert r.costs_modelled
    assert r.is_quotable


def test_the_legacy_not_implemented_exception_still_imports():
    """A caller that imported this to signal 'backtesting is unavailable' must not break."""
    from app.engines.backtesting.replay import PointInTimeReplayNotImplemented

    assert issubclass(PointInTimeReplayNotImplemented, NotImplementedError)


ALL_TESTS = [v for k, v in sorted(globals().items()) if k.startswith("test_")]

if __name__ == "__main__":
    passed = failed = 0
    for t in ALL_TESTS:
        try:
            t()
            print(f"PASS  {t.__name__}")
            passed += 1
        except Exception as exc:  # noqa: BLE001
            print(f"FAIL  {t.__name__}: {exc}")
            failed += 1
    print(f"\n{passed}/{passed + failed} passed")
    raise SystemExit(1 if failed else 0)
