from __future__ import annotations

from dataclasses import dataclass

from adjustments.models import Account, RunMembership

VIEWPORTS = ("phone", "tablet")


@dataclass(frozen=True)
class Node:
    role: str
    label: str = ""
    value: str | None = None
    ref: str | None = None
    collapsed: bool = False
    children: tuple["Node", ...] = ()


def _correction(record, run_id: str | None) -> Node:
    child = (
        Node("settlement", "Settled by statement run", run_id, run_id),
    ) if run_id else (
        Node(
            "pending",
            "Correction pending the next run",
            "pending",
            record.adjustment_id,
        ),
    )
    return Node(
        "correction",
        "Correction",
        str(record.delta),
        record.adjustment_id,
        children=child
        + (
            Node(
                "link",
                "Correction detail",
                record.record_id,
                record.adjustment_id,
            ),
        ),
    )


def build_statement(
    account: Account, period_key: str, viewport: str = "phone"
) -> Node:
    if viewport not in VIEWPORTS:
        raise ValueError("unknown viewport")
    invoice = account.invoices.get(
        period_key=period_key, revision_of__isnull=True
    )
    latest = (
        account.statement_runs.filter(status="issued")
        .order_by("-run_number")
        .first()
    )
    settled = []
    pending = []
    for record in account.records.filter(original_invoice=invoice).order_by(
        "acceptance_ordinal", "pk"
    ):
        membership = (
            RunMembership.objects.filter(
                record=record, run__status="issued"
            )
            .select_related("run")
            .order_by("run__run_number")
            .first()
        )
        node = _correction(
            record, membership.run.run_id if membership else None
        )
        (settled if membership else pending).append(node)
    children = [
        Node("header", f"Billing period {invoice.period_key}"),
        Node(
            "group",
            "Charges",
            children=tuple(
                Node(
                    "line",
                    f"{item.description} ({item.quantity})",
                    str(item.total),
                )
                for item in invoice.line_items.order_by("pk")
            ),
        ),
        Node(
            "issued_amount",
            "Invoice total as issued",
            str(
                invoice.issued_total
                if invoice.issued_total is not None
                else invoice.total
            ),
        ),
    ]
    if latest is not None:
        children.append(
            Node(
                "run",
                f"As of statement run {latest.run_number}",
                str(latest.recorded_demand),
                latest.run_id,
            )
        )
    else:
        children.append(Node("run", "No statement run issued", "0.00"))
    if settled:
        children.append(
            Node(
                "group",
                "Corrections settled by statement runs",
                children=tuple(settled),
            )
        )
    if pending:
        children.append(
            Node(
                "group",
                "Corrections pending the next run",
                children=tuple(pending),
            )
        )
    children.append(Node("amount_due", "Amount due", str(account.amount_due)))
    return Node(
        "statement", f"Statement {invoice.invoice_id}", children=tuple(children)
    )


def _render(node: Node) -> str:
    attrs = f' data-role="{node.role}"'
    if node.ref is not None:
        attrs += f' data-ref="{node.ref}"'
    if node.collapsed:
        attrs += ' data-collapsed="true"'
    label = f'<span class="label">{node.label}</span>' if node.label else ""
    value = (
        f'<span class="value">{node.value}</span>'
        if node.value is not None
        else ""
    )
    return (
        f"<div{attrs}>{label}{value}"
        f"{''.join(_render(child) for child in node.children)}</div>"
    )


def render_html(node: Node, viewport: str = "phone") -> str:
    if viewport not in VIEWPORTS:
        raise ValueError("unknown viewport")
    return (
        f'<html><body class="statement viewport-{viewport}">'
        f"{_render(node)}</body></html>"
    )
