"""
Pull selected fields from PartKeepr and merge into parts_inventory parts.

Fields merged into parts_inventory.parts:
  - average_price
  - minimum_stock_level
  - comment

Matching strategy (in order):
  1) source_partkeepr_id
  2) storage location name + part_number
  3) storage location name + part name

Usage example:
  python scripts/merge_partkeepr_part_fields.py ^
    --source-host 192.168.50.5 --source-port 3306 ^
    --source-db PartKeepr --source-user partkeepr --source-password partkeepr ^
    --target-host 192.168.50.5 --target-port 3306 ^
    --target-db parts_inventory --target-user root --target-password secret ^
    --dry-run
"""
from __future__ import annotations

import argparse
from collections import defaultdict
from typing import Dict, List, Optional, Tuple

import pymysql
from pymysql.cursors import DictCursor


def connect(host: str, port: int, user: str, password: str, database: str):
    return pymysql.connect(
        host=host,
        port=port,
        user=user,
        password=password,
        database=database,
        cursorclass=DictCursor,
        autocommit=False,
    )


def parse_args():
    parser = argparse.ArgumentParser(
        description="Merge average price, minimum stock level, and comment from PartKeepr into parts_inventory."
    )
    parser.add_argument("--source-host", required=True)
    parser.add_argument("--source-port", type=int, default=3306)
    parser.add_argument("--source-db", required=True)
    parser.add_argument("--source-user", required=True)
    parser.add_argument("--source-password", required=True)
    parser.add_argument("--target-host", required=True)
    parser.add_argument("--target-port", type=int, default=3306)
    parser.add_argument("--target-db", required=True)
    parser.add_argument("--target-user", required=True)
    parser.add_argument("--target-password", required=True)
    parser.add_argument("--dry-run", action="store_true")
    return parser.parse_args()


def get_table_columns(conn, table_name: str) -> set[str]:
    with conn.cursor() as cur:
        cur.execute(
            """
            SELECT COLUMN_NAME
            FROM information_schema.COLUMNS
            WHERE TABLE_SCHEMA = DATABASE()
              AND TABLE_NAME = %s
            """,
            (table_name,),
        )
        return {r["COLUMN_NAME"] for r in cur.fetchall()}


def choose_column(available: set[str], candidates: List[str]) -> Optional[str]:
    for c in candidates:
        if c in available:
            return c
    return None


def ensure_target_columns(dst_conn) -> None:
    with dst_conn.cursor() as cur:
        try:
            cur.execute("ALTER TABLE parts ADD COLUMN average_price DECIMAL(12,4) NOT NULL DEFAULT 0 AFTER unit")
        except pymysql.OperationalError:
            pass
        try:
            cur.execute("ALTER TABLE parts ADD COLUMN minimum_stock_level INT NOT NULL DEFAULT 0 AFTER average_price")
        except pymysql.OperationalError:
            pass
        try:
            cur.execute("ALTER TABLE parts ADD COLUMN comment TEXT NULL AFTER minimum_stock_level")
        except pymysql.OperationalError:
            pass
    dst_conn.commit()


def load_source_partkeepr_rows(src_conn) -> Tuple[List[dict], str, str]:
    part_columns = get_table_columns(src_conn, "Part")
    avg_col = choose_column(part_columns, ["averagePrice", "average_price", "price"])
    min_col = choose_column(part_columns, ["minStockLevel", "minimumStockLevel", "minimum_stock_level", "min_stock_level"])
    if not avg_col:
        raise RuntimeError("PartKeepr Part table has no average price column candidate.")
    if not min_col:
        raise RuntimeError("PartKeepr Part table has no minimum stock level column candidate.")

    sql = f"""
        SELECT
          p.id AS source_partkeepr_id,
          p.name,
          p.internalPartNumber AS part_number,
          p.comment,
          p.`{avg_col}` AS average_price,
          p.`{min_col}` AS minimum_stock_level,
          s.name AS storage_location_name
        FROM Part p
        LEFT JOIN StorageLocation s ON p.storageLocation_id = s.id
        ORDER BY p.id
    """
    with src_conn.cursor() as cur:
        cur.execute(sql)
        rows = cur.fetchall()
    return rows, avg_col, min_col


def load_target_rows(dst_conn) -> List[dict]:
    with dst_conn.cursor() as cur:
        cur.execute(
            """
            SELECT
              p.id,
              p.source_partkeepr_id,
              p.name,
              p.part_number,
              s.name AS storage_location_name
            FROM parts p
            LEFT JOIN storage_locations s ON s.id = p.storage_location_id
            ORDER BY p.id
            """
        )
        return cur.fetchall()


def normalize_text(value) -> str:
    if value is None:
        return ""
    return str(value).strip()


def normalize_number(value, default: float = 0.0) -> float:
    try:
        if value is None:
            return default
        return float(value)
    except (TypeError, ValueError):
        return default


def choose_target_row(src_row: dict, by_source_id: Dict[int, dict], by_loc_pn: Dict[Tuple[str, str], List[dict]], by_loc_name: Dict[Tuple[str, str], List[dict]]) -> Tuple[Optional[dict], str]:
    src_id = src_row.get("source_partkeepr_id")
    if src_id in by_source_id:
        return by_source_id[src_id], "source_partkeepr_id"

    loc_name = normalize_text(src_row.get("storage_location_name"))
    part_number = normalize_text(src_row.get("part_number"))
    name = normalize_text(src_row.get("name"))

    if loc_name and part_number:
        candidates = by_loc_pn.get((loc_name, part_number), [])
        if len(candidates) == 1:
            return candidates[0], "location+part_number"

    if loc_name and name:
        candidates = by_loc_name.get((loc_name, name), [])
        if len(candidates) == 1:
            return candidates[0], "location+name"

    return None, "unmatched"


def main():
    args = parse_args()
    src_conn = connect(args.source_host, args.source_port, args.source_user, args.source_password, args.source_db)
    dst_conn = connect(args.target_host, args.target_port, args.target_user, args.target_password, args.target_db)

    try:
        ensure_target_columns(dst_conn)
        source_rows, avg_col, min_col = load_source_partkeepr_rows(src_conn)
        target_rows = load_target_rows(dst_conn)

        by_source_id: Dict[int, dict] = {}
        by_loc_pn: Dict[Tuple[str, str], List[dict]] = defaultdict(list)
        by_loc_name: Dict[Tuple[str, str], List[dict]] = defaultdict(list)

        for row in target_rows:
            if row.get("source_partkeepr_id") is not None:
                by_source_id[int(row["source_partkeepr_id"])] = row
            loc_name = normalize_text(row.get("storage_location_name"))
            part_number = normalize_text(row.get("part_number"))
            name = normalize_text(row.get("name"))
            if loc_name and part_number:
                by_loc_pn[(loc_name, part_number)].append(row)
            if loc_name and name:
                by_loc_name[(loc_name, name)].append(row)

        updated = 0
        skipped_unmatched = 0
        skipped_ambiguous = 0
        matched_by = defaultdict(int)

        with dst_conn.cursor() as cur:
            for src_row in source_rows:
                dst_row, method = choose_target_row(src_row, by_source_id, by_loc_pn, by_loc_name)
                if not dst_row:
                    # Try to distinguish no-match vs ambiguous match.
                    loc_name = normalize_text(src_row.get("storage_location_name"))
                    part_number = normalize_text(src_row.get("part_number"))
                    name = normalize_text(src_row.get("name"))
                    ambiguous = False
                    if loc_name and part_number and len(by_loc_pn.get((loc_name, part_number), [])) > 1:
                        ambiguous = True
                    if loc_name and name and len(by_loc_name.get((loc_name, name), [])) > 1:
                        ambiguous = True
                    if ambiguous:
                        skipped_ambiguous += 1
                    else:
                        skipped_unmatched += 1
                    continue

                matched_by[method] += 1
                avg_price = normalize_number(src_row.get("average_price"), 0.0)
                min_stock = int(normalize_number(src_row.get("minimum_stock_level"), 0))
                comment = normalize_text(src_row.get("comment"))

                cur.execute(
                    """
                    UPDATE parts
                    SET average_price = %s,
                        minimum_stock_level = %s,
                        comment = %s,
                        updated_at = CURRENT_TIMESTAMP
                    WHERE id = %s
                    """,
                    (avg_price, min_stock, comment, dst_row["id"]),
                )
                updated += 1

        if args.dry_run:
            dst_conn.rollback()
        else:
            dst_conn.commit()

        print(
            "Merge complete:",
            f"source_rows={len(source_rows)}",
            f"updated={updated}",
            f"matched_by_source_id={matched_by['source_partkeepr_id']}",
            f"matched_by_location_part_number={matched_by['location+part_number']}",
            f"matched_by_location_name={matched_by['location+name']}",
            f"skipped_unmatched={skipped_unmatched}",
            f"skipped_ambiguous={skipped_ambiguous}",
            f"source_avg_col={avg_col}",
            f"source_min_col={min_col}",
            f"dry_run={args.dry_run}",
        )
    except Exception:
        dst_conn.rollback()
        raise
    finally:
        src_conn.close()
        dst_conn.close()


if __name__ == "__main__":
    main()
