"""
Peer-group universe construction (StockLab final engineering pass, Part B2).

This is the module `app/workers/recompute.py` and `app/engines/valuation/multiples.py` both
referred to as "not yet implemented" — `recompute_security()`'s docstring even names this exact
file path. Its absence was not a missing nice-to-have: because nothing supplied
`peer_metric_values`, every percentile in the scoring engine was computed against an empty peer
list, `percentile_rank()` returned its neutral 50.0, and **every security's every pillar score
came out at exactly 50**. See docs/AUDIT_PEER_GROUPS_B2.md.

Everything here is the DB half of the job — loading rows and handing them to the pure,
unit-tested logic in `app/engines/scoring/peer_groups.py`. No scoring decision is made in this
file, deliberately: it imports `sqlalchemy` and therefore cannot be executed in the build
environment, so it holds as little judgment as possible.
"""
from __future__ import annotations

from typing import Optional

from sqlalchemy import select
from sqlalchemy.orm import Session

from app.core.config import get_settings
from app.core.logging import get_logger
from app.engines.scoring.peer_groups import (
    PeerMetricRow,
    build_peer_metric_values,
    industry_median_multiples,
)
from app.models import Company, Industry, Metric, Price, Security, Shares
from app.models.financials import FinancialPeriod
from app.workers.celery_app import celery_app

logger = get_logger(__name__)

#: Same thresholds as build_snapshot_from_db()'s bucketing, in millions of the reporting currency.
#: Duplicated deliberately rather than imported from recompute.py, which would create an import
#: cycle (recompute imports this module). Kept in sync by test_peer_group_buckets_match_snapshot.
_BUCKET_THRESHOLDS = (("LARGE", 10_000.0), ("MID", 2_000.0), ("SMALL", 300.0))


def market_cap_bucket_for(market_cap: Optional[float]) -> Optional[str]:
    if market_cap is None:
        return None
    for name, floor in _BUCKET_THRESHOLDS:
        if market_cap >= floor:
            return name
    return "MICRO"


def _market_caps_by_security(db: Session) -> dict[str, float]:
    """Latest close price x latest diluted share count, per security.

    Market cap is not persisted anywhere in this schema — it is computed at snapshot time from
    `price * shares` — so the cap bucket a peer group needs has to be rebuilt here. Two batched
    queries, not per-security lookups: this runs over the whole universe.
    """
    prices: dict[str, float] = {}
    for security_id, close in db.execute(
        select(Price.security_id, Price.close).order_by(Price.security_id, Price.date.desc())
    ).all():
        prices.setdefault(security_id, close)

    shares: dict[str, float] = {}
    rows = db.execute(
        select(FinancialPeriod.security_id, Shares.diluted_shares, Shares.shares_outstanding)
        .join(Shares, Shares.financial_period_id == FinancialPeriod.id)
        .where(FinancialPeriod.period_type == "FY")
        .order_by(FinancialPeriod.security_id, FinancialPeriod.period_end.desc())
    ).all()
    for security_id, diluted, outstanding in rows:
        count = diluted or outstanding
        if count:
            shares.setdefault(security_id, count)

    return {
        sid: prices[sid] * shares[sid]
        for sid in prices.keys() & shares.keys()
    }


def load_metric_universe(db: Session) -> list[PeerMetricRow]:
    """Every security's current metric values, classified for peer grouping.

    One query for the metric rows plus the two batched queries behind the cap buckets — not a
    per-security lookup, which at N securities x M metrics would be exactly the N+1 pattern
    Part A2 spent its time removing.
    """
    caps = _market_caps_by_security(db)
    rows = db.execute(
        select(
            Metric.security_id, Metric.metric_key, Metric.value,
            Metric.industry_id, Industry.sector_id,
        )
        .join(Security, Security.id == Metric.security_id)
        .join(Company, Company.id == Security.company_id)
        .join(Industry, Industry.id == Metric.industry_id)
        .where(Metric.value.is_not(None))
    ).all()
    return [
        PeerMetricRow(
            security_id=security_id,
            metric_key=metric_key,
            value=value,
            industry_id=industry_id,
            sector_id=sector_id,
            market_cap_bucket=market_cap_bucket_for(caps.get(security_id)),
        )
        for security_id, metric_key, value, industry_id, sector_id in rows
    ]


def peer_metric_values_for(
    universe: list[PeerMetricRow], security_id: str,
) -> dict[str, dict[str, list[float]]]:
    """The `peer_metric_values` dict `recompute_security()` takes, for one security.

    The target's own classification is read out of the universe rather than passed in separately,
    so a security's peers are always grouped by exactly the classification its own metric rows
    carry — there is no way for the two to disagree.
    """
    own = next((r for r in universe if r.security_id == security_id), None)
    if own is None:
        logger.info("peer_groups.security_not_in_universe", security_id=security_id)
        return {}
    return build_peer_metric_values(
        universe, security_id, own.industry_id, own.sector_id, own.market_cap_bucket,
    ).as_dict()


REFERENCE_MULTIPLE_KEYS = ("pe", "forward_pe", "ev_to_ebitda", "p_fcf", "ev_to_fcf")


def industry_medians_for(
    all_industry_medians: dict, universe: list[PeerMetricRow], security_id: str,
) -> Optional[dict]:
    """The `{metric_key: median}` dict for ONE security's own industry, or None."""
    own = next((r for r in universe if r.security_id == security_id), None)
    if own is None or own.industry_id is None:
        return None
    return all_industry_medians.get(own.industry_id)


def industry_reference_multiples_from_universe(
    universe: list[PeerMetricRow], min_group_size: int = 5,
) -> dict[str, dict[str, float]]:
    """Same as `industry_reference_multiples()` but reuses an already-loaded universe."""
    return industry_median_multiples(universe, REFERENCE_MULTIPLE_KEYS, min_group_size=min_group_size)


def industry_reference_multiples(db: Session, min_group_size: int = 5) -> dict[str, dict[str, float]]:
    """`{industry_id: {metric_key: median}}` — what `ReferenceSource.INDUSTRY_MEDIAN` needs.

    Industries with fewer than `min_group_size` companies reporting a given multiple are absent
    from the result entirely, so the caller keeps its existing fallback to the self-historical
    reference instead of receiving a two-company "industry median".
    """
    return industry_median_multiples(
        load_metric_universe(db), REFERENCE_MULTIPLE_KEYS, min_group_size=min_group_size,
    )


@celery_app.task(name="app.workers.peer_groups.recompute_universe_task")
def recompute_universe_task(passes: int = 2) -> int:
    """Recompute every security with a real peer universe, loading the universe ONCE per pass.

    **Why two passes by default.** The peer universe is built from `metrics` rows, and
    `recompute_security()` is what writes `metrics` rows. On a cold database — or after a schema
    change, or a first ingestion — pass 1 therefore runs against an empty or stale universe and
    produces scores with no peer comparison (which, since Part B2, are honestly reported as
    INSUFFICIENT_DATA rather than a neutral 50). Pass 2 reloads the now-populated universe and
    produces the real percentile scores. On a warm database pass 1 already has a full universe and
    pass 2 only refines it with the freshly-written values.

    This is a deliberate cost trade, not an oversight: the alternative is splitting
    `recompute_security()` into a metrics phase and a scoring phase so the universe can be built
    between them, which is the better architecture and a real refactor of a module that cannot be
    executed or regression-tested in this build environment. Recorded as the recommended next step
    in docs/AUDIT_PEER_GROUPS_B2.md rather than attempted blind.

    Recomputing securities one at a time via `recompute_security_task` also works and now gets
    real peers, but rebuilds the universe per call — O(N) universe loads for N securities. Use
    this task after a universe-wide ingestion.
    """
    from app.core.db import SessionLocal
    from app.workers.recompute import recompute_security

    db = SessionLocal()
    count = 0
    try:
        for pass_number in range(1, max(1, passes) + 1):
            universe = load_metric_universe(db)
            security_ids = sorted({row.security_id for row in universe})
            if not security_ids:
                # Cold start: no metrics exist yet, so fall back to every security we know about.
                security_ids = sorted(db.execute(select(Security.id)).scalars().all())
                logger.warning("peer_groups.universe_empty_using_all_securities",
                               pass_number=pass_number, securities=len(security_ids))
            settings = get_settings()
            all_medians = industry_reference_multiples_from_universe(
                universe, min_group_size=settings.INDUSTRY_MULTIPLE_MIN_GROUP_SIZE,
            )
            for security_id in security_ids:
                recompute_security(
                    db, security_id,
                    peer_metric_values=peer_metric_values_for(universe, security_id),
                    industry_medians=industry_medians_for(all_medians, universe, security_id),
                )
            count = len(security_ids)
            logger.info("peer_groups.pass_complete", pass_number=pass_number, securities=count)
        return count
    finally:
        db.close()
