"""Entity resolution — alias normalization and merchant linking."""

from __future__ import annotations

import re
import uuid
from datetime import datetime, timezone
from typing import Any

from app.database import get_connection

_SUFFIX_RE = re.compile(
    r"(?i)\b(pvt\.?|private|ltd\.?|limited|llc|inc\.?|corp\.?|company|co\.?|"
    r"services?|india|marketplace|retail|seller)\b"
)


def canonicalize_name(name: str) -> str:
    t = (name or "").lower()
    t = t.replace("&", " and ")
    t = _SUFFIX_RE.sub(" ", t)
    t = re.sub(r"[^a-z0-9]+", " ", t)
    t = re.sub(r"\s+", " ", t).strip()
    return t


# Seed aliases for common Indian merchants / platforms
SEED_ALIASES: dict[str, list[str]] = {
    "amazon": ["amazon", "amazon india", "amazon pay", "amzn", "amazon seller services"],
    "zepto": ["zepto", "zepto marketplace"],
    "blinkit": ["blinkit", "grofers"],
    "swiggy": ["swiggy", "swiggy instamart"],
    "zomato": ["zomato", "zomato ltd", "zomato limited"],
    "myntra": ["myntra"],
    "flipkart": ["flipkart"],
    "uber": ["uber", "uber india"],
    "ola": ["ola", "ani technologies"],
    "jio": ["jio", "reliance jio"],
    "airtel": ["airtel", "bharti airtel"],
}


def ensure_seed_aliases(org_id: str) -> int:
    """Upsert curated merchant alias clusters for an org."""
    now = datetime.now(timezone.utc).isoformat()
    created = 0
    with get_connection() as conn:
        for canonical, aliases in SEED_ALIASES.items():
            row = conn.execute(
                """
                SELECT e.id FROM fdi_entities e
                JOIN fdi_entity_aliases a ON a.entity_id=e.id
                WHERE a.org_id=? AND a.alias=?
                LIMIT 1
                """,
                (org_id, canonical),
            ).fetchone()
            if row:
                entity_id = str(row["id"] if isinstance(row, dict) else row[0])
            else:
                entity_id = str(uuid.uuid4())
                conn.execute(
                    """
                    INSERT INTO fdi_entities(id, org_id, entity_type, canonical_name, attributes_json, created_at)
                    VALUES (?, ?, 'merchant', ?, ?, ?)
                    """,
                    (entity_id, org_id, canonical.title(), None, now),
                )
                created += 1
            for alias in aliases:
                exists = conn.execute(
                    "SELECT id FROM fdi_entity_aliases WHERE org_id=? AND alias=? LIMIT 1",
                    (org_id, alias),
                ).fetchone()
                if exists:
                    continue
                conn.execute(
                    """
                    INSERT INTO fdi_entity_aliases(id, org_id, entity_id, alias, source, weight, created_at)
                    VALUES (?, ?, ?, ?, 'seed', 1.0, ?)
                    """,
                    (str(uuid.uuid4()), org_id, entity_id, alias, now),
                )
                created += 1
    return created


def resolve_entity_id(conn: Any, org_id: str, name: str) -> str | None:
    alias = canonicalize_name(name)
    if not alias:
        return None
    # exact alias
    row = conn.execute(
        """
        SELECT entity_id FROM fdi_entity_aliases
        WHERE org_id=? AND alias=?
        LIMIT 1
        """,
        (org_id, alias),
    ).fetchone()
    if row:
        return str(row["entity_id"] if isinstance(row, dict) else row[0])
    # token containment against seed-like aliases
    row = conn.execute(
        """
        SELECT entity_id, alias FROM fdi_entity_aliases
        WHERE org_id=? AND length(alias) >= 4
        """,
        (org_id,),
    ).fetchall()
    for r in row:
        a = str(r["alias"] if isinstance(r, dict) else r[1])
        if a in alias or alias in a:
            return str(r["entity_id"] if isinstance(r, dict) else r[0])
    return None


def relink_line_counterparties(org_id: str, document_id: str | None = None) -> int:
    """Re-resolve counterparty_entity_id using alias table."""
    ensure_seed_aliases(org_id)
    updated = 0
    with get_connection() as conn:
        if document_id:
            rows = conn.execute(
                """
                SELECT id, description_norm FROM fdi_line_items
                WHERE org_id=? AND document_id=? AND description_norm IS NOT NULL
                """,
                (org_id, document_id),
            ).fetchall()
        else:
            rows = conn.execute(
                """
                SELECT id, description_norm FROM fdi_line_items
                WHERE org_id=? AND description_norm IS NOT NULL
                LIMIT 5000
                """,
                (org_id,),
            ).fetchall()
        for r in rows:
            lid = str(r["id"] if isinstance(r, dict) else r[0])
            desc = str(r["description_norm"] if isinstance(r, dict) else r[1])
            eid = resolve_entity_id(conn, org_id, desc)
            if not eid:
                continue
            conn.execute(
                "UPDATE fdi_line_items SET counterparty_entity_id=? WHERE id=?",
                (eid, lid),
            )
            updated += 1
    return updated
