"""Matching, Merge, Duplikatauflösung und Adress-Fallback."""

from __future__ import annotations

from collections import defaultdict
from typing import Callable

from app.mapping import FieldKey
from app.models import CustomerRecord
from app.normalize import (
    clean_text,
    format_zip,
    is_empty,
    normalize_name,
    normalize_plate,
    normalize_vin,
)

ADDRESS_FIELDS = (FieldKey.STREET, FieldKey.ZIP, FieldKey.CITY)


def plate_key(record: CustomerRecord) -> str:
    return normalize_plate(record.get(FieldKey.PLATE))


def vin_key(record: CustomerRecord) -> str:
    return normalize_vin(record.get(FieldKey.VIN))


def address_no_key(record: CustomerRecord) -> str:
    return clean_text(record.get(FieldKey.ADDRESS_NO))


def customer_no_key(record: CustomerRecord) -> str:
    return clean_text(record.get(FieldKey.CUSTOMER_NO))


def company_key(record: CustomerRecord) -> str:
    return normalize_name(
        record.get(FieldKey.COMPANY)
        or record.get(FieldKey.CUSTOMER_NAME)
        or ""
    )


def person_name_key(record: CustomerRecord) -> str:
    first = normalize_name(record.get(FieldKey.FIRST_NAME))
    last = normalize_name(record.get(FieldKey.LAST_NAME))
    if not first or not last:
        return ""
    return f"{last}|{first}"


def customer_name_key(record: CustomerRecord) -> str:
    company = company_key(record)
    if company:
        return company
    last = normalize_name(record.get(FieldKey.LAST_NAME))
    return last


def address_score(record: CustomerRecord) -> int:
    return sum(1 for key in ADDRESS_FIELDS if not is_empty(record.get(key)))


def is_address_complete(record: CustomerRecord) -> bool:
    return address_score(record) == 3


def record_completeness(record: CustomerRecord) -> int:
    return sum(1 for value in record.values.values() if not is_empty(value))


def address_tuple(record: CustomerRecord) -> tuple[str, str, str]:
    return (
        normalize_name(record.get(FieldKey.STREET)),
        format_zip(record.get(FieldKey.ZIP)) if not is_empty(record.get(FieldKey.ZIP)) else "",
        normalize_name(record.get(FieldKey.CITY)),
    )


def merge_fill(target: CustomerRecord, source: CustomerRecord) -> None:
    """Füllt nur leere/Platzhalterfelder in target aus source."""
    for key, value in source.values.items():
        if is_empty(target.get(key)) and not is_empty(value):
            target.set(key, value)


def _merge_compatible(records: list[CustomerRecord]) -> CustomerRecord:
    ranked = sorted(
        records,
        key=lambda rec: (address_score(rec), record_completeness(rec), -rec.row_index),
        reverse=True,
    )
    merged = ranked[0].clone()
    for other in ranked[1:]:
        merge_fill(merged, other)
    merged.source = "A"
    return merged


def _group_conflicts(records: list[CustomerRecord]) -> list[str]:
    issues: list[str] = []
    plates = {plate_key(r) for r in records if plate_key(r)}
    vins = {vin_key(r) for r in records if vin_key(r)}
    if len(plates) > 1:
        issues.append("Widersprüchliche Kennzeichen bei gleichem Fahrzeugbezug")
    if len(vins) > 1:
        issues.append("Widersprüchliche Fahrgestellnummern bei gleichem Fahrzeugbezug")

    complete_addresses = {address_tuple(r) for r in records if is_address_complete(r)}
    if len(complete_addresses) > 1:
        issues.append("Widersprüchliche vollständige Adressen")
    return issues


def _connected_vehicle_groups(records: list[CustomerRecord]) -> list[list[int]]:
    parent = list(range(len(records)))

    def find(i: int) -> int:
        while parent[i] != i:
            parent[i] = parent[parent[i]]
            i = parent[i]
        return i

    def union(a: int, b: int) -> None:
        ra, rb = find(a), find(b)
        if ra != rb:
            parent[rb] = ra

    plate_index: dict[str, int] = {}
    vin_index: dict[str, int] = {}
    for i, record in enumerate(records):
        plate = plate_key(record)
        vin = vin_key(record)
        if plate:
            if plate in plate_index:
                union(i, plate_index[plate])
            else:
                plate_index[plate] = i
        if vin:
            if vin in vin_index:
                union(i, vin_index[vin])
            else:
                vin_index[vin] = i

    buckets: dict[int, list[int]] = defaultdict(list)
    for i in range(len(records)):
        buckets[find(i)].append(i)
    return list(buckets.values())


def deduplicate_records(records: list[CustomerRecord]) -> tuple[list[CustomerRecord], list[CustomerRecord]]:
    """Fasst Datensätze mit gleichem Kennzeichen oder gleicher VIN zusammen."""
    if not records:
        return [], []

    unique: list[CustomerRecord] = []
    conflicts: list[CustomerRecord] = []

    for group_idx in _connected_vehicle_groups(records):
        group = [records[i] for i in group_idx]
        if len(group) == 1:
            unique.append(group[0])
            continue

        issues = _group_conflicts(group)
        if issues:
            conflict = _merge_compatible(group)
            conflict.conflict = True
            conflict.issues.extend(issues)
            conflict.match_status = "Konflikt"
            conflicts.append(conflict)
            continue

        plates = {plate_key(r) for r in group if plate_key(r)}
        vins = {vin_key(r) for r in group if vin_key(r)}
        same_plate_and_vin = len(plates) <= 1 and len(vins) <= 1 and (plates or vins)
        same_address_no = len({address_no_key(r) for r in group if address_no_key(r)}) <= 1

        if same_plate_and_vin or (is_address_complete(max(group, key=address_score)) and same_address_no):
            merged = _merge_compatible(group)
            unique.append(merged)
            continue

        best_addr = max(group, key=lambda r: (address_score(r), record_completeness(r)))
        others_strictly_worse = all(
            r is best_addr
            or address_score(r) < address_score(best_addr)
            or (
                address_score(r) == address_score(best_addr)
                and record_completeness(r) < record_completeness(best_addr)
            )
            for r in group
        )
        if others_strictly_worse and address_score(best_addr) == 3:
            unique.append(_merge_compatible(group))
            continue

        conflict = _merge_compatible(group)
        conflict.conflict = True
        conflict.issues.append("Duplikat nicht eindeutig auflösbar")
        conflict.match_status = "Konflikt"
        conflicts.append(conflict)

    return unique, conflicts


MatchKeyFn = Callable[[CustomerRecord], object]

MATCH_LEVELS: list[tuple[str, MatchKeyFn]] = [
    (
        "Kennzeichen + VIN",
        lambda r: (plate_key(r), vin_key(r)) if plate_key(r) and vin_key(r) else None,
    ),
    ("Fahrgestellnummer", lambda r: vin_key(r) or None),
    ("Kennzeichen", lambda r: plate_key(r) or None),
    ("Adressnummer", lambda r: address_no_key(r) or None),
    ("Kundennummer", lambda r: customer_no_key(r) or None),
    ("Kundenname", lambda r: customer_name_key(r) or None),
    ("Vorname + Nachname", lambda r: person_name_key(r) or None),
]


def _index_by(records: list[CustomerRecord], keyfn: MatchKeyFn) -> dict[object, list[CustomerRecord]]:
    index: dict[object, list[CustomerRecord]] = defaultdict(list)
    for record in records:
        key = keyfn(record)
        if key:
            index[key].append(record)
    return index


def _choose_candidate(candidates: list[CustomerRecord]) -> tuple[CustomerRecord | None, bool]:
    if not candidates:
        return None, False
    if len(candidates) == 1:
        return candidates[0], False

    best_addr_score = max(address_score(c) for c in candidates)
    top = [c for c in candidates if address_score(c) == best_addr_score]
    if len(top) == 1:
        return top[0], False

    addresses = {address_tuple(c) for c in top}
    if len(addresses) == 1:
        best_complete = max(top, key=record_completeness)
        return best_complete, False

    return None, True


def match_and_enrich(
    records: list[CustomerRecord],
    templates: list[CustomerRecord],
) -> list[CustomerRecord]:
    """Ordnet A-Datensätze Beispielzeilen aus B zu und füllt nur leere Felder."""
    indexes = [(label, _index_by(templates, keyfn), keyfn) for label, keyfn in MATCH_LEVELS]

    for record in records:
        matched = False
        for label, index, keyfn in indexes:
            key = keyfn(record)
            if not key:
                continue
            candidates = index.get(key, [])
            if not candidates:
                continue
            chosen, conflict = _choose_candidate(candidates)
            if conflict:
                record.conflict = True
                record.issues.append(f"Nicht eindeutig zuordenbar ({label})")
                record.match_status = "Konflikt"
                matched = True
                break
            if chosen is None:
                continue
            merge_fill(record, chosen)
            record.match_status = label
            matched = True
            break
        if not matched and record.match_status != "Konflikt":
            record.match_status = "Kein Match (Neukunde)"
    return records


ADDRESS_FALLBACK_LEVELS: list[tuple[str, MatchKeyFn]] = [
    ("Adressnummer", lambda r: address_no_key(r) or None),
    ("Kennzeichen", lambda r: plate_key(r) or None),
    ("Fahrgestellnummer", lambda r: vin_key(r) or None),
    ("Kundennummer", lambda r: customer_no_key(r) or None),
    ("Firmenname", lambda r: company_key(r) or None),
    ("Nachname + Vorname", lambda r: person_name_key(r) or None),
]


def fill_missing_addresses(
    records: list[CustomerRecord],
    lookup_pool: list[CustomerRecord],
) -> int:
    filled = 0
    complete_pool = [r for r in lookup_pool if is_address_complete(r)]
    complete_indexes = [
        (label, _index_by(complete_pool, keyfn), keyfn) for label, keyfn in ADDRESS_FALLBACK_LEVELS
    ]

    for record in records:
        if is_address_complete(record):
            continue
        for _label, index, keyfn in complete_indexes:
            key = keyfn(record)
            if not key:
                continue
            candidates = [c for c in index.get(key, []) if c is not record]
            if not candidates:
                continue
            chosen = max(candidates, key=lambda r: (address_score(r), record_completeness(r)))
            for field_name in ADDRESS_FIELDS + (FieldKey.COUNTRY,):
                if is_empty(record.get(field_name)) and not is_empty(chosen.get(field_name)):
                    record.set(field_name, chosen.get(field_name))
            filled += 1
            break
    return filled
