from collections import defaultdict, deque
from dataclasses import dataclass
from decimal import Decimal

from .models import UnitConversion, UnitOfMeasurement


@dataclass(frozen=True)
class _UnitMeta:
    name: str
    code: str
    dimension: str


class UnitConversionNormalizer:
    """
    Converts arbitrary unit quantities to a deterministic canonical unit per
    conversion-connected component, without changing API response structure.
    """

    def __init__(self):
        units = UnitOfMeasurement.objects.only("id", "name", "code", "dimension")
        self._units = {
            unit.id: _UnitMeta(
                name=unit.name,
                code=unit.code,
                dimension=(unit.dimension or "").strip(),
            )
            for unit in units
        }
        self._canonical_unit_id_by_unit_id = {}
        self._factor_to_canonical_by_unit_id = {}
        self._build_maps()

    def _build_maps(self):
        graph = defaultdict(list)
        outgoing = defaultdict(int)
        incoming = defaultdict(int)

        conversions = UnitConversion.objects.values_list(
            "from_unit_id", "to_unit_id", "factor"
        )
        for from_unit_id, to_unit_id, factor in conversions:
            from_meta = self._units.get(from_unit_id)
            to_meta = self._units.get(to_unit_id)
            if not from_meta or not to_meta:
                continue
            if not from_meta.dimension or from_meta.dimension != to_meta.dimension:
                continue

            decimal_factor = Decimal(str(factor))
            if decimal_factor <= 0:
                continue

            outgoing[from_unit_id] += 1
            incoming[to_unit_id] += 1
            graph[from_unit_id].append((to_unit_id, decimal_factor))
            graph[to_unit_id].append((from_unit_id, Decimal("1") / decimal_factor))

        visited = set()
        for unit_id in self._units:
            if unit_id in visited:
                continue

            component = []
            queue = deque([unit_id])
            visited.add(unit_id)

            while queue:
                current = queue.popleft()
                component.append(current)
                for neighbor_id, _ in graph.get(current, []):
                    if neighbor_id in visited:
                        continue
                    visited.add(neighbor_id)
                    queue.append(neighbor_id)

            canonical_unit_id = self._pick_canonical_unit(
                component, outgoing, incoming
            )
            factor_map = self._factor_map_to_canonical(canonical_unit_id, graph)

            for member_id in component:
                self._canonical_unit_id_by_unit_id[member_id] = canonical_unit_id
                self._factor_to_canonical_by_unit_id[member_id] = factor_map.get(
                    member_id, Decimal("1")
                )

    def _pick_canonical_unit(self, component, outgoing, incoming):
        def _score(unit_id):
            meta = self._units[unit_id]
            out_count = outgoing.get(unit_id, 0)
            in_count = incoming.get(unit_id, 0)
            return (
                -(out_count - in_count),
                -out_count,
                meta.code or "",
                meta.name or "",
                str(unit_id),
            )

        return min(component, key=_score)

    @staticmethod
    def _factor_map_to_canonical(canonical_unit_id, graph):
        factor_to_canonical = {canonical_unit_id: Decimal("1")}
        queue = deque([canonical_unit_id])

        while queue:
            current = queue.popleft()
            current_factor = factor_to_canonical[current]
            for neighbor_id, current_to_neighbor_factor in graph.get(current, []):
                neighbor_factor = current_factor / current_to_neighbor_factor
                if neighbor_id in factor_to_canonical:
                    continue
                factor_to_canonical[neighbor_id] = neighbor_factor
                queue.append(neighbor_id)

        return factor_to_canonical

    def normalize_quantity(self, quantity, unit_id, fallback_unit_name=""):
        quantity_decimal = Decimal(str(quantity or 0))
        if not unit_id or unit_id not in self._units:
            return quantity_decimal, fallback_unit_name

        canonical_unit_id = self._canonical_unit_id_by_unit_id.get(unit_id, unit_id)
        factor = self._factor_to_canonical_by_unit_id.get(unit_id, Decimal("1"))
        canonical_meta = self._units.get(canonical_unit_id)
        canonical_name = canonical_meta.name if canonical_meta else fallback_unit_name

        return quantity_decimal * factor, canonical_name
