"""
Merge duplicate parts where one record has part_number like PK-<id>
and another record has blank part_number for the same name.

Keeps the blank-part-number row as canonical (usually has correct unit like 'pcs'),
moves images from PK row, copies source_partkeepr_id, then deletes PK row.
"""
from __future__ import annotations

import argparse

import pymysql
from pymysql.cursors import DictCursor


def parse_args():
    parser = argparse.ArgumentParser(description="Reconcile PK-* duplicate parts in parts_inventory.")
    parser.add_argument("--host", required=True)
    parser.add_argument("--port", type=int, default=3306)
    parser.add_argument("--database", default="parts_inventory")
    parser.add_argument("--user", required=True)
    parser.add_argument("--password", required=True)
    parser.add_argument("--dry-run", action="store_true")
    return parser.parse_args()


def main():
    args = parse_args()
    conn = pymysql.connect(
        host=args.host,
        port=args.port,
        user=args.user,
        password=args.password,
        database=args.database,
        cursorclass=DictCursor,
        autocommit=False,
    )
    merged = 0
    moved_images = 0
    try:
        with conn.cursor() as cur:
            cur.execute(
                """
                SELECT
                  blank.id AS keep_id,
                  blank.name,
                  blank.part_number AS keep_part_number,
                  blank.quantity AS keep_quantity,
                  blank.unit AS keep_unit,
                  blank.source_partkeepr_id AS keep_source_id,
                  pk.id AS drop_id,
                  pk.part_number AS drop_part_number,
                  pk.quantity AS drop_quantity,
                  pk.unit AS drop_unit,
                  pk.source_partkeepr_id AS drop_source_id
                FROM parts blank
                JOIN parts pk ON pk.name = blank.name
                WHERE (blank.part_number IS NULL OR blank.part_number = '')
                  AND pk.part_number LIKE 'PK-%'
                  AND blank.id <> pk.id
                ORDER BY blank.id, pk.id
                """
            )
            pairs = cur.fetchall()

            for pair in pairs:
                keep_id = pair["keep_id"]
                drop_id = pair["drop_id"]

                # Repoint images to canonical row.
                cur.execute("UPDATE part_images SET part_id = %s WHERE part_id = %s", (keep_id, drop_id))
                moved_images += cur.rowcount

                # Merge important fields: keep canonical unit/qty unless missing.
                merged_quantity = pair["keep_quantity"] if pair["keep_quantity"] not in (None, 0) else pair["drop_quantity"]
                merged_unit = pair["keep_unit"] if pair["keep_unit"] else (pair["drop_unit"] or "")
                merged_source = pair["keep_source_id"] if pair["keep_source_id"] is not None else pair["drop_source_id"]

                # Avoid unique-key collision on parts.source_partkeepr_id while both rows still exist.
                if merged_source is not None:
                    cur.execute("UPDATE parts SET source_partkeepr_id = NULL WHERE id = %s", (drop_id,))
                    cur.execute(
                        """
                        SELECT id FROM parts
                        WHERE source_partkeepr_id = %s
                          AND id NOT IN (%s, %s)
                        LIMIT 1
                        """,
                        (merged_source, keep_id, drop_id),
                    )
                    collision = cur.fetchone()
                    if collision:
                        # Keep merge proceeding, but do not move source id when another row already owns it.
                        merged_source = pair["keep_source_id"]

                cur.execute(
                    """
                    UPDATE parts
                    SET quantity = %s,
                        unit = %s,
                        source_partkeepr_id = %s,
                        updated_at = CURRENT_TIMESTAMP
                    WHERE id = %s
                    """,
                    (merged_quantity, merged_unit, merged_source, keep_id),
                )

                # Remove duplicate PK row.
                cur.execute("DELETE FROM parts WHERE id = %s", (drop_id,))
                merged += 1

        if args.dry_run:
            conn.rollback()
        else:
            conn.commit()
        print(f"Reconcile complete: merged={merged} moved_images={moved_images} dry_run={args.dry_run}")
    except Exception:
        conn.rollback()
        raise
    finally:
        conn.close()


if __name__ == "__main__":
    main()
