"""Central role-aware portfolio aggregation for verified PostFinance projections.

Official account totals are reconciliation controls, never additive holdings. Source FX
rates and line values can be rounded independently, so component CHF values are allocated
deterministically to the signed-off source subtotal without inventing an extra position.
"""

from __future__ import annotations

import json
from dataclasses import dataclass
from decimal import ROUND_HALF_UP, Decimal
from sqlite3 import Connection

MONEY = Decimal("0.01")


@dataclass(frozen=True)
class OfficialCashComponent:
    account_id: str
    currency: str
    amount_original: Decimal
    source_fx_rate_to_chf: Decimal
    amount_chf: Decimal
    snapshot_date: str
    valuation_at: str
    reconciliation_method: str = "official_total_pro_rata_rounded_fx_v1"


@dataclass(frozen=True)
class OfficialPositionOverride:
    account_id: str
    instrument_id: str
    asset_class: str
    quantity: Decimal
    market_price_original: Decimal
    price_currency: str
    market_value_chf: Decimal
    snapshot_date: str
    valuation_at: str
    reconciliation_method: str = "official_asset_class_residual_v1"


@dataclass(frozen=True)
class AuditedMarketValue:
    account_id: str
    instrument_id: str
    value_chf: Decimal
    close: Decimal | None
    currency: str
    as_of: str
    completed_at: str
    provider: str
    quantity: Decimal | None = None


def _decimal(value: object) -> Decimal:
    return Decimal(str(value))


def _money(value: Decimal) -> Decimal:
    return value.quantize(MONEY, rounding=ROUND_HALF_UP)


def postfinance_account_roles(conn: Connection) -> dict[str, str]:
    """Return canonical PostFinance role mappings without name-based inference."""

    return {
        str(row["role"]): str(row["account_id"])
        for row in conn.execute(
            "SELECT role,account_id FROM postfinance_account_roles ORDER BY role"
        ).fetchall()
    }


def _allocate_pro_rata(values: list[Decimal], target: Decimal) -> list[Decimal]:
    """Reconcile rounded source components to an authoritative signed-off subtotal."""

    if not values:
        return []
    raw_total = sum(values, Decimal("0"))
    same_nonzero_sign = raw_total * target > 0
    allocated = [
        _money(value * target / raw_total if same_nonzero_sign else value)
        for value in values
    ]
    residual = _money(target - sum(allocated, Decimal("0")))
    if residual:
        largest = max(range(len(values)), key=lambda index: abs(values[index]))
        allocated[largest] = _money(allocated[largest] + residual)
    if sum(allocated, Decimal("0")) != _money(target):
        raise ValueError("Official portfolio component allocation does not reconcile")
    return allocated


def latest_official_postfinance_cash(
    conn: Connection,
    *,
    as_of: str | None = None,
    data_cutoff: str | None = None,
) -> list[OfficialCashComponent]:
    roles = postfinance_account_roles(conn)
    account_id = roles.get("etrading_cash")
    if not account_id:
        return []
    conditions: list[str] = []
    params: list[str] = []
    if as_of:
        conditions.append("substr(snapshot_at,1,10)<=substr(?,1,10)")
        params.append(as_of)
    if data_cutoff:
        conditions.append("datetime(created_at)<=datetime(?)")
        params.append(data_cutoff)
    where = f"WHERE {' AND '.join(conditions)}" if conditions else ""
    snapshot = conn.execute(
        f"""SELECT snapshot_id,snapshot_at,cash_chf
            FROM postfinance_snapshots {where}
            ORDER BY snapshot_at DESC,datetime(created_at) DESC LIMIT 1""",
        params,
    ).fetchone()
    if not snapshot:
        return []
    rows = conn.execute(
        """SELECT currency,amount_original,fx_rate_to_chf
           FROM postfinance_snapshot_cash WHERE snapshot_id=? ORDER BY currency""",
        (snapshot["snapshot_id"],),
    ).fetchall()
    raw_values = [
        _decimal(row["amount_original"]) * _decimal(row["fx_rate_to_chf"])
        for row in rows
    ]
    target = _money(_decimal(snapshot["cash_chf"]))
    allocated = _allocate_pro_rata(raw_values, target)
    valuation_at = str(snapshot["snapshot_at"])
    return [
        OfficialCashComponent(
            account_id=account_id,
            currency=str(row["currency"]),
            amount_original=_decimal(row["amount_original"]),
            source_fx_rate_to_chf=_decimal(row["fx_rate_to_chf"]),
            amount_chf=value,
            snapshot_date=valuation_at[:10],
            valuation_at=valuation_at,
        )
        for row, value in zip(rows, allocated, strict=True)
    ]


def latest_official_postfinance_positions(
    conn: Connection,
    *,
    as_of: str | None = None,
    data_cutoff: str | None = None,
) -> dict[tuple[str, str], OfficialPositionOverride]:
    roles = postfinance_account_roles(conn)
    account_id = roles.get("etrading_depot")
    if not account_id:
        return {}
    conditions: list[str] = []
    params: list[str] = []
    if as_of:
        conditions.append("substr(snapshot_at,1,10)<=substr(?,1,10)")
        params.append(as_of)
    if data_cutoff:
        conditions.append("datetime(created_at)<=datetime(?)")
        params.append(data_cutoff)
    where = f"WHERE {' AND '.join(conditions)}" if conditions else ""
    snapshot = conn.execute(
        f"""SELECT snapshot_id,snapshot_at,stocks_chf,etfs_chf
            FROM postfinance_snapshots {where}
            ORDER BY snapshot_at DESC,datetime(created_at) DESC LIMIT 1""",
        params,
    ).fetchone()
    if not snapshot:
        return {}
    rows = conn.execute(
        """SELECT instrument_id,asset_class,quantity,price_original,price_currency,value_chf
           FROM postfinance_snapshot_positions
           WHERE snapshot_id=? ORDER BY asset_class,instrument_id""",
        (snapshot["snapshot_id"],),
    ).fetchall()
    targets = {
        "stock": _money(_decimal(snapshot["stocks_chf"])),
        "etf": _money(_decimal(snapshot["etfs_chf"])),
    }
    allocated_by_instrument: dict[str, Decimal] = {}
    for asset_class, target in targets.items():
        members = [row for row in rows if str(row["asset_class"]).lower() == asset_class]
        allocated = _allocate_pro_rata(
            [_decimal(row["value_chf"]) for row in members], target
        )
        allocated_by_instrument.update(
            {
                str(row["instrument_id"]): value
                for row, value in zip(members, allocated, strict=True)
            }
        )
    valuation_at = str(snapshot["snapshot_at"])
    return {
        (account_id, str(row["instrument_id"])): OfficialPositionOverride(
            account_id=account_id,
            instrument_id=str(row["instrument_id"]),
            asset_class=str(row["asset_class"]).lower(),
            quantity=_decimal(row["quantity"]),
            market_price_original=_decimal(row["price_original"]),
            price_currency=str(row["price_currency"]),
            market_value_chf=allocated_by_instrument[str(row["instrument_id"])],
            snapshot_date=valuation_at[:10],
            valuation_at=valuation_at,
        )
        for row in rows
    }


def post_snapshot_postfinance_quantity_deltas(
    conn: Connection,
    positions: dict[tuple[str, str], OfficialPositionOverride],
    *,
    as_of: str | None = None,
) -> dict[str, Decimal]:
    """Replay confirmed buys/sells after the official snapshot baseline."""

    if not positions:
        return {}
    roles = postfinance_account_roles(conn)
    account_ids = [
        account_id
        for role in ("etrading_depot", "etrading_cash")
        if (account_id := roles.get(role))
    ]
    if not account_ids:
        return {}
    snapshot_date = min(item.snapshot_date for item in positions.values())
    effective_as_of = as_of or "9999-12-31"
    placeholders = ",".join("?" for _ in account_ids)
    rows = conn.execute(
        f"""SELECT instrument_id,lower(transaction_type) transaction_type,quantity
            FROM transactions
            WHERE account_id IN ({placeholders})
              AND instrument_id IS NOT NULL
              AND COALESCE(is_confirmed,0)=1
              AND COALESCE(is_voided,0)=0
              AND lower(transaction_type) IN ('buy','sell','partial_sell','full_sell')
              AND trade_date>? AND trade_date<=?""",
        (*account_ids, snapshot_date, effective_as_of),
    ).fetchall()
    deltas: dict[str, Decimal] = {}
    for row in rows:
        instrument_id = str(row["instrument_id"])
        quantity = abs(_decimal(row["quantity"]))
        if str(row["transaction_type"]) in {"sell", "partial_sell", "full_sell"}:
            quantity = -quantity
        deltas[instrument_id] = deltas.get(instrument_id, Decimal("0")) + quantity
    return deltas


def latest_audited_market_values(
    conn: Connection,
) -> dict[tuple[str, str], AuditedMarketValue]:
    """Return the latest immutable market-analysis values without writing on GET."""

    row = conn.execute(
        """SELECT pas.as_of,pas.summary_json,mdr.completed_at
           FROM portfolio_analysis_snapshots pas
           JOIN market_data_runs mdr ON mdr.run_id=pas.run_id
           ORDER BY pas.as_of DESC,pas.created_at DESC LIMIT 1"""
    ).fetchone()
    if not row:
        return {}
    try:
        payload = json.loads(row["summary_json"] or "{}")
    except json.JSONDecodeError:
        return {}
    values: dict[tuple[str, str], AuditedMarketValue] = {}
    for item in payload.get("positions", []):
        if not isinstance(item, dict) or item.get("quality_status") != "fresh":
            continue
        account_id = str(item.get("account_id") or "")
        instrument_id = str(item.get("instrument_id") or "")
        if not account_id or not instrument_id or item.get("value_chf") in {None, ""}:
            continue
        provenance = item.get("price_input_provenance")
        provider = (
            str(item.get("provider") or "")
            or (str(provenance.get("provider") or "") if isinstance(provenance, dict) else "")
        )
        values[(account_id, instrument_id)] = AuditedMarketValue(
            account_id=account_id,
            instrument_id=instrument_id,
            value_chf=_decimal(item["value_chf"]),
            close=_decimal(item["close"]) if item.get("close") not in {None, ""} else None,
            currency=str(item.get("currency") or ""),
            as_of=str(row["as_of"]),
            completed_at=str(row["completed_at"]),
            provider=provider,
            quantity=(
                _decimal(item["quantity"])
                if item.get("quantity") not in {None, ""}
                else None
            ),
        )
    return values
