"""
Point-in-time replay engine (StockLab final engineering pass, Part C).

**Status change: this file was previously a documented NOT IMPLEMENTED placeholder** — see
`docs/AUDIT_BACKTESTING.md` and the `PointInTimeReplayNotImplemented` exception kept at the bottom
of this module for compatibility. It now contains a real, executed, unit-tested replay loop.

**What that does and does not mean.** The ENGINE is real: it walks a portfolio forward through
rebalance dates, asks a caller-supplied data source what was knowable at each date, refuses inputs
that would introduce look-ahead bias, and produces a period-return series for
`compute_performance_report()`. **No backtest of StockLab's actual strategy has been run**, because
this sandbox has no database and no historical price data — see `docs/AUDIT_BACKTESTING_C.md` §5.
An engine that has never been fed real data has produced no results, and none are claimed.

## The four ways a backtest lies, and what this engine does about each

1. **Look-ahead bias** — using information that did not exist at the decision date. This is the
   one the engine can and does enforce mechanically: `_assert_no_lookahead()` checks every
   returned observation's `known_as_of` against the rebalance date and raises
   `LookAheadBiasError` if any is later. It is an assertion, not a filter: silently dropping a
   future-dated row would hide a broken data source. `app/engines/metrics/point_in_time.py::
   as_of_snapshot()` is the primitive a caller's data source should use to build the snapshot in
   the first place; this check is the backstop for when it doesn't.

2. **Survivorship bias** — running the backtest over today's surviving universe. The engine
   cannot detect this from the data, so it does the next best thing: `universe_at(date)` is a
   REQUIRED callable, and if it returns the identical membership for every rebalance date the
   result carries `survivorship_bias_suspected=True` and a warning string. A constant universe is
   not proof of bias (a genuinely static index exists) and a varying one is not proof of its
   absence, so this is reported as a suspicion, never as a verdict.

3. **Unadjusted prices** — a stock split silently registering as a −50% return. `PriceObservation`
   carries an explicit `is_split_adjusted` flag with NO default: a caller must state it. When any
   observation is unadjusted the result carries `unadjusted_prices_used=True`. The engine does not
   guess, and does not adjust anything itself — it has no corporate-action data.

4. **Frictionless trading** — no commissions, no spread, no market impact. `cost_per_unit_turnover`
   is applied to every rebalance, defaults to `0.0`, and when it IS zero the result says
   `costs_modelled=False` so a headline CAGR is never mistaken for a net-of-costs figure.

## What this engine deliberately does NOT do

- **It does not fetch anything.** `SnapshotSource` is a protocol the caller implements against a
  real database. That keeps the loop pure, dependency-free and genuinely testable here, and it
  keeps the look-ahead enforcement in one place regardless of where data comes from.
- **It does not decide what to buy.** `decide` is a caller-supplied function from an observation
  to a target weight. Running StockLab's own live scoring pipeline through it is the intended use
  and is what makes the backtest meaningful, but wiring that in requires the DB-backed snapshot
  source, so it is a worker-level job, not an engine-level one.
- **It does not model dividends, taxes, borrowing, shorting, or partial fills.** Returns are price
  returns of a long-only, fully-invested, weight-target portfolio. Stated, not implied.
"""
from __future__ import annotations

from dataclasses import dataclass, field
from datetime import date
from typing import Callable, Iterable, Optional, Protocol, Sequence


class PointInTimeReplayNotImplemented(NotImplementedError):
    """Kept for compatibility with the pre-Part-C placeholder API.

    No longer raised by this module. A caller that imported it to signal "backtesting is not
    available" should now call `run_replay()`; this class remains only so that import does not
    break.
    """


class LookAheadBiasError(AssertionError):
    """Raised when a data source returns an observation dated after the decision date.

    Deliberately an error, not a filter. A data source that hands back future information is
    broken, and quietly discarding the row would produce a backtest that runs cleanly on top of a
    bug — which is the exact failure mode this whole module exists to prevent.
    """


@dataclass(frozen=True)
class PriceObservation:
    """One security's state at one rebalance date, as a point-in-time observer would have seen it.

    `known_as_of` is the latest date any information in this observation became public — the
    filing date of the most recent financial period used, or the price date, whichever is later.
    The caller computes it; the engine checks it.
    """

    security_id: str
    price: float
    known_as_of: date
    is_split_adjusted: bool          # no default: the caller must state this
    payload: Optional[object] = None  # e.g. a FinancialSnapshot, for the decide() function


class SnapshotSource(Protocol):
    """What the replay loop needs from a data source. Implemented against the DB by a worker."""

    def universe_at(self, as_of: date) -> Sequence[str]:
        """Securities that were investable on `as_of` — NOT today's survivors."""

    def observe(self, security_id: str, as_of: date) -> Optional[PriceObservation]:
        """Point-in-time observation, or None if the security had no usable data on that date."""


#: A decision function: given an observation, return a desired portfolio weight in [0, 1].
#: Returning 0.0 means "do not hold". The engine normalises the returned weights.
DecideFn = Callable[[PriceObservation], float]


@dataclass
class RebalanceRecord:
    as_of: date
    holdings: dict[str, float]           # security_id -> weight, post-rebalance
    turnover: float                      # sum of |weight change|, 0..2
    period_return: Optional[float]       # return realised BETWEEN the previous date and this one
    securities_considered: int
    securities_priced: int


@dataclass
class ReplayResult:
    rebalances: list[RebalanceRecord] = field(default_factory=list)
    period_returns: list[float] = field(default_factory=list)
    turnover_per_period: list[float] = field(default_factory=list)
    #: Integrity flags — every one of these must be read before a result is quoted.
    survivorship_bias_suspected: bool = False
    unadjusted_prices_used: bool = False
    costs_modelled: bool = False
    warnings: list[str] = field(default_factory=list)

    @property
    def periods(self) -> int:
        return len(self.period_returns)

    @property
    def is_quotable(self) -> bool:
        """True only when no integrity flag is raised and there is at least one period.

        A backtest with a raised flag is not necessarily wrong, but it is not a number to publish
        without saying which flag is raised and why. Consumers should refuse to render a headline
        CAGR when this is False.
        """
        return (
            self.periods > 0
            and not self.survivorship_bias_suspected
            and not self.unadjusted_prices_used
            and self.costs_modelled
        )


def _assert_no_lookahead(obs: PriceObservation, as_of: date) -> None:
    if obs.known_as_of > as_of:
        raise LookAheadBiasError(
            f"{obs.security_id}: observation known_as_of={obs.known_as_of} is AFTER the rebalance "
            f"date {as_of}. A point-in-time observer could not have seen this."
        )


def _normalise_weights(raw: dict[str, float]) -> dict[str, float]:
    """Scale non-negative weights to sum to 1. An all-zero book stays empty (cash), not 1/N."""
    positive = {k: v for k, v in raw.items() if v and v > 0}
    total = sum(positive.values())
    if total <= 0:
        return {}
    return {k: v / total for k, v in positive.items()}


def _portfolio_return(
    holdings: dict[str, float], prev_prices: dict[str, float], new_prices: dict[str, float],
) -> tuple[float, list[str]]:
    """Weighted price return of `holdings` between two price snapshots.

    A holding whose price is missing at the later date is treated as a 0% return for that period
    AND reported by name, rather than dropped: dropping it would quietly re-weight the portfolio
    onto its survivors, which is survivorship bias reintroduced at the smallest scale.
    """
    total = 0.0
    dropped: list[str] = []
    for security_id, weight in holdings.items():
        p0 = prev_prices.get(security_id)
        p1 = new_prices.get(security_id)
        if p0 is None or p1 is None or p0 <= 0:
            dropped.append(security_id)
            continue
        total += weight * (p1 / p0 - 1.0)
    return total, dropped


def run_replay(
    source: SnapshotSource,
    decide: DecideFn,
    rebalance_dates: Iterable[date],
    cost_per_unit_turnover: float = 0.0,
) -> ReplayResult:
    """Walk a long-only portfolio forward through `rebalance_dates`.

    At each date: ask `source` for the investable universe, observe each member point-in-time,
    assert no look-ahead, mark the previous book to the new prices to realise a period return,
    then ask `decide` for new target weights and record the turnover cost.
    """
    dates = sorted(set(rebalance_dates))
    result = ReplayResult()
    if len(dates) < 2:
        result.warnings.append(
            f"insufficient_rebalance_dates:{len(dates)} — at least 2 are needed to realise a "
            f"single period return"
        )
        return result

    universes: list[tuple[str, ...]] = []
    holdings: dict[str, float] = {}
    prev_prices: dict[str, float] = {}

    for index, as_of in enumerate(dates):
        universe = tuple(sorted(source.universe_at(as_of)))
        universes.append(universe)

        observations: dict[str, PriceObservation] = {}
        for security_id in universe:
            obs = source.observe(security_id, as_of)
            if obs is None:
                continue
            _assert_no_lookahead(obs, as_of)
            if not obs.is_split_adjusted:
                result.unadjusted_prices_used = True
            observations[security_id] = obs

        new_prices = {sid: obs.price for sid, obs in observations.items()}

        period_return: Optional[float] = None
        if index > 0:
            gross, dropped = _portfolio_return(holdings, prev_prices, new_prices)
            if dropped:
                result.warnings.append(
                    f"{as_of}: no price for {len(dropped)} held securities "
                    f"({', '.join(sorted(dropped)[:5])}{'...' if len(dropped) > 5 else ''}) — "
                    f"held at 0% for this period rather than dropped"
                )
            period_return = gross

        targets = _normalise_weights({sid: decide(obs) for sid, obs in observations.items()})
        turnover = sum(
            abs(targets.get(sid, 0.0) - holdings.get(sid, 0.0))
            for sid in set(targets) | set(holdings)
        )

        if period_return is not None:
            # Costs are charged at the rebalance that CAUSES the turnover, deducted from the
            # period return being realised at that same date.
            period_return -= turnover * cost_per_unit_turnover
            result.period_returns.append(period_return)
            result.turnover_per_period.append(turnover)

        result.rebalances.append(RebalanceRecord(
            as_of=as_of, holdings=dict(targets), turnover=turnover, period_return=period_return,
            securities_considered=len(universe), securities_priced=len(observations),
        ))
        holdings = targets
        prev_prices = new_prices

    if len(set(universes)) == 1:
        result.survivorship_bias_suspected = True
        result.warnings.append(
            "universe_constant_across_all_rebalance_dates — the same securities were investable "
            "on every date. This is what a survivorship-biased backtest looks like (today's "
            "survivors replayed into the past). It is not proof of bias: a genuinely static "
            "universe would look identical. Confirm the data source has real historical listing "
            "membership before quoting these results."
        )

    result.costs_modelled = cost_per_unit_turnover > 0
    if not result.costs_modelled:
        result.warnings.append(
            "cost_per_unit_turnover=0 — returns are gross of commissions, spread and market "
            "impact, and are NOT comparable to a real portfolio's net return."
        )
    return result
