#!/usr/bin/env python3
# ============================================================
# ECENG BOTTLE RECOGNITION
#
# Alur:
# 1. Ambil project ec_project yang is_active='Y'
# 2. Ambil ec_capture status_process='NEW' ORDER BY id ASC
# 3. Baca image dari capture/
# 4. YOLO mendeteksi botol
# 5. Simpan image hasil ke recognized/
# 6. UPDATE ec_capture.qty/status_process/remark
# 7. Hapus image asli hanya jika proses sukses
#
# MODE AWAL:
# - Model YOLO 1 class: "bottle"
# - Hasil qty = jumlah botol
#
# MODE 16 JENIS:
# - Ganti model dengan model 16 class
# - ec_bottle_config.class_id digunakan untuk validasi jenis
# ============================================================

import os
import sys
import time
import shutil
import traceback
from pathlib import Path
from datetime import datetime

import cv2
import pymysql
from ultralytics import YOLO


# ============================================================
# CONFIG
# ============================================================

BASE_DIR = Path("/home/horek940/portalmkits/api/eceng")

CAPTURE_DIR = BASE_DIR / "capture"
RECOGNIZED_DIR = BASE_DIR / "recognized"

# Untuk tahap pertama gunakan model 1 class "bottle".
# Setelah model 16 jenis selesai dilatih, ganti ke:
# MODEL_PATH = BASE_DIR / "model" / "bottle16.pt"
MODEL_PATH = Path(
    os.getenv(
        "ECENG_MODEL_PATH",
        str(BASE_DIR / "model" / "bottle.pt")
    )
)

# Confidence minimal deteksi.
DEFAULT_CONFIDENCE = float(
    os.getenv("ECENG_CONFIDENCE", "0.70")
)

# Berapa detik menunggu sebelum cek database lagi jika tidak ada NEW.
POLL_INTERVAL = float(
    os.getenv("ECENG_POLL_INTERVAL", "1.0")
)

# Batas jumlah botol yang dianggap valid.
# Kebutuhan saat ini 0 / 1 / 2.
MAX_BOTTLE_QTY = int(
    os.getenv("ECENG_MAX_BOTTLE_QTY", "2")
)

# Jika True, semua hasil di atas MAX_BOTTLE_QTY menjadi ERROR.
LIMIT_QTY = os.getenv("ECENG_LIMIT_QTY", "Y").upper() == "Y"

# Validasi jenis botol:
# N = hanya counting, tidak memaksa kecocokan product_id.
# Y = jika ec_bottle_config memiliki konfigurasi untuk product aktif,
#     class yang ditemukan harus sesuai.
STRICT_PRODUCT_TYPE = (
    os.getenv("ECENG_STRICT_PRODUCT_TYPE", "N").upper() == "Y"
)

# Nama file output mengikuti nama_file capture.
# Jika True, overwrite file recognized yang sudah ada.
OVERWRITE_RECOGNIZED = True


# ============================================================
# DATABASE
# Jangan hardcode password produksi.
# Set environment variable:
# ECENG_DB_HOST
# ECENG_DB_PORT
# ECENG_DB_NAME
# ECENG_DB_USER
# ECENG_DB_PASSWORD
# ============================================================

DB_CONFIG = {
    "host": os.getenv("ECENG_DB_HOST", "127.0.0.1"),
    "port": int(os.getenv("ECENG_DB_PORT", "3306")),
    "user": os.getenv("ECENG_DB_USER", "root"),
    "password": os.getenv("ECENG_DB_PASSWORD", ""),
    "database": os.getenv("ECENG_DB_NAME", "portalmkits"),
    "charset": "utf8mb4",
    "autocommit": False,
    "cursorclass": pymysql.cursors.DictCursor,
}


# ============================================================
# LOG
# ============================================================

def log(message):
    print(
        f"[{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] "
        f"{message}",
        flush=True
    )


# ============================================================
# DATABASE
# ============================================================

def get_connection():
    return pymysql.connect(**DB_CONFIG)


def get_active_project(conn):
    """
    Ambil project aktif.
    Jika ada lebih dari satu Y, project dengan id terbesar
    dianggap project aktif terakhir.
    """
    sql = """
        SELECT
            id,
            project_code,
            product_barcode,
            product_id,
            product_nama,
            preform_barcode,
            preform_id,
            preform_nama
        FROM ec_project
        WHERE is_active = 'Y'
        ORDER BY id DESC
        LIMIT 1
    """

    with conn.cursor() as cur:
        cur.execute(sql)
        return cur.fetchone()


def get_new_capture(conn, project_code):
    """
    Ambil 1 capture paling lama yang masih NEW.
    ORDER BY id ASC.
    """
    sql = """
        SELECT
            id,
            project_code,
            date_at,
            code,
            name_file,
            qty,
            status_process,
            status_warehouse,
            remark
        FROM ec_capture
        WHERE project_code = %s
          AND status_process = 'NEW'
        ORDER BY id ASC
        LIMIT 1
    """

    with conn.cursor() as cur:
        cur.execute(sql, (project_code,))
        return cur.fetchone()


def get_bottle_config(conn, product_id):
    """
    Cari konfigurasi botol berdasarkan product_id.
    """
    if not product_id:
        return None

    sql = """
        SELECT
            id,
            bottle_code,
            bottle_name,
            product_id,
            product_barcode,
            preform_id,
            preform_barcode,
            class_id,
            confidence_min,
            is_active
        FROM ec_bottle_config
        WHERE product_id = %s
          AND is_active = 'Y'
        ORDER BY id ASC
        LIMIT 1
    """

    with conn.cursor() as cur:
        cur.execute(sql, (product_id,))
        return cur.fetchone()


def update_capture_success(
    conn,
    capture_id,
    qty,
    remark
):
    sql = """
        UPDATE ec_capture
        SET
            qty = %s,
            status_process = 'FINISH',
            remark = %s,
            updated_at = NOW()
        WHERE id = %s
    """

    with conn.cursor() as cur:
        cur.execute(sql, (qty, remark, capture_id))

    conn.commit()


def update_capture_error(
    conn,
    capture_id,
    remark
):
    sql = """
        UPDATE ec_capture
        SET
            status_process = 'ERROR',
            remark = %s,
            updated_at = NOW()
        WHERE id = %s
    """

    with conn.cursor() as cur:
        cur.execute(sql, (remark, capture_id))

    conn.commit()


# ============================================================
# IMAGE
# ============================================================

def get_capture_path(name_file):
    return CAPTURE_DIR / name_file


def get_recognized_path(name_file):
    return RECOGNIZED_DIR / name_file


def draw_detection(
    image,
    box,
    class_name,
    confidence,
    index
):
    x1, y1, x2, y2 = box

    x1 = max(0, int(x1))
    y1 = max(0, int(y1))
    x2 = max(0, int(x2))
    y2 = max(0, int(y2))

    cv2.rectangle(
        image,
        (x1, y1),
        (x2, y2),
        (0, 255, 0),
        2
    )

    label = (
        f"{index}. {class_name} "
        f"{confidence:.2f}"
    )

    text_y = y1 - 10
    if text_y < 20:
        text_y = y1 + 25

    cv2.putText(
        image,
        label,
        (x1, text_y),
        cv2.FONT_HERSHEY_SIMPLEX,
        0.65,
        (0, 255, 0),
        2,
        cv2.LINE_AA
    )


def recognize_image(model, image_path):
    """
    Return:
        image_result
        detections

    detections:
    [
        {
            "class_id": 0,
            "class_name": "bottle",
            "confidence": 0.95,
            "box": [x1,y1,x2,y2]
        }
    ]
    """

    image = cv2.imread(str(image_path))

    if image is None:
        raise RuntimeError(
            f"Gambar tidak dapat dibaca: {image_path}"
        )

    results = model.predict(
        source=image,
        conf=DEFAULT_CONFIDENCE,
        verbose=False
    )

    detections = []

    if not results:
        return image, detections

    result = results[0]

    names = result.names

    if result.boxes is None:
        return image, detections

    for box in result.boxes:
        class_id = int(box.cls[0].item())
        confidence = float(box.conf[0].item())

        coords = box.xyxy[0].tolist()

        if isinstance(names, dict):
            class_name = str(
                names.get(class_id, f"class_{class_id}")
            )
        else:
            class_name = str(names[class_id])

        detections.append({
            "class_id": class_id,
            "class_name": class_name,
            "confidence": confidence,
            "box": coords,
        })

    # Urutkan dari kiri ke kanan agar nomor botol stabil.
    detections.sort(
        key=lambda d: (
            (d["box"][0] + d["box"][2]) / 2
        )
    )

    for index, detection in enumerate(
        detections,
        start=1
    ):
        draw_detection(
            image=image,
            box=detection["box"],
            class_name=detection["class_name"],
            confidence=detection["confidence"],
            index=index
        )

    # Tampilkan jumlah pada gambar.
    qty = len(detections)

    height, width = image.shape[:2]

    count_text = f"COUNT = {qty}"

    cv2.rectangle(
        image,
        (10, 10),
        (250, 65),
        (0, 0, 0),
        -1
    )

    cv2.putText(
        image,
        count_text,
        (20, 50),
        cv2.FONT_HERSHEY_SIMPLEX,
        1.0,
        (0, 255, 255),
        2,
        cv2.LINE_AA
    )

    return image, detections


def save_recognized_image(image, output_path):
    output_path.parent.mkdir(
        parents=True,
        exist_ok=True
    )

    if output_path.exists() and not OVERWRITE_RECOGNIZED:
        raise RuntimeError(
            f"File recognized sudah ada: {output_path}"
        )

    ok = cv2.imwrite(
        str(output_path),
        image
    )

    if not ok:
        raise RuntimeError(
            f"Gagal menyimpan recognized image: {output_path}"
        )


# ============================================================
# VALIDATION
# ============================================================

def validate_detections(
    detections,
    project,
    bottle_config
):
    """
    Menghasilkan:
        valid: bool
        remark: str

    Pada mode awal STRICT_PRODUCT_TYPE=N,
    fungsi ini hanya memastikan jumlah <= batas.

    Pada mode STRICT_PRODUCT_TYPE=Y,
    class_id harus sama dengan ec_bottle_config.class_id.
    """

    qty = len(detections)

    if LIMIT_QTY and qty > MAX_BOTTLE_QTY:
        return (
            False,
            f"Jumlah botol {qty} melebihi batas "
            f"{MAX_BOTTLE_QTY}"
        )

    # Tidak ada konfigurasi product.
    if bottle_config is None:
        if STRICT_PRODUCT_TYPE:
            return (
                False,
                "Konfigurasi botol tidak ditemukan "
                f"untuk product_id={project.get('product_id')}"
            )

        return (
            True,
            f"{qty} bottle detected; "
            "product config belum ditemukan"
        )

    expected_class_id = int(
        bottle_config["class_id"]
    )

    wrong_classes = []

    for detection in detections:
        if detection["class_id"] != expected_class_id:
            wrong_classes.append(
                f'{detection["class_name"]}'
                f'(class={detection["class_id"]})'
            )

    if wrong_classes:
        return (
            False,
            "Wrong bottle type. "
            f"Expected {bottle_config['bottle_code']} "
            f"class={expected_class_id}; "
            f"detected={','.join(wrong_classes)}"
        )

    return (
        True,
        f"{bottle_config['bottle_code']} "
        f"qty={qty}"
    )


# ============================================================
# PROCESS ONE CAPTURE
# ============================================================

def process_capture(model, project, capture):
    capture_id = capture["id"]
    name_file = capture["name_file"]

    capture_path = get_capture_path(name_file)
    recognized_path = get_recognized_path(name_file)

    log(
        f"PROCESS capture id={capture_id} "
        f"file={name_file}"
    )

    if not capture_path.exists():
        raise FileNotFoundError(
            f"File capture tidak ditemukan: {capture_path}"
        )

    conn = get_connection()

    try:
        bottle_config = get_bottle_config(
            conn,
            project.get("product_id")
        )

        image_result, detections = recognize_image(
            model,
            capture_path
        )

        qty = len(detections)

        valid, validation_remark = validate_detections(
            detections,
            project,
            bottle_config
        )

        detection_detail = ", ".join(
            [
                f'{d["class_name"]}:{d["confidence"]:.2f}'
                for d in detections
            ]
        )

        if not detection_detail:
            detection_detail = "none"

        remark = (
            f"qty={qty}; "
            f"detections={detection_detail}; "
            f"{validation_remark}"
        )

        # Simpan gambar hasil SEBELUM update FINISH.
        save_recognized_image(
            image_result,
            recognized_path
        )

        if not valid:
            update_capture_error(
                conn,
                capture_id,
                remark
            )

            # ERROR tetap menyimpan capture asli
            # agar dapat diperiksa ulang.
            log(
                f"ERROR capture id={capture_id}: "
                f"{remark}"
            )

            return False

        update_capture_success(
            conn,
            capture_id,
            qty,
            remark
        )

        # Hapus capture hanya setelah:
        # 1. recognition sukses
        # 2. recognized image sukses
        # 3. database FINISH sukses
        capture_path.unlink()

        log(
            f"FINISH capture id={capture_id} "
            f"qty={qty} "
            f"recognized={recognized_path}"
        )

        return True

    except Exception:
        conn.rollback()
        raise

    finally:
        conn.close()


# ============================================================
# MAIN LOOP
# ============================================================

def check_directories():
    CAPTURE_DIR.mkdir(
        parents=True,
        exist_ok=True
    )

    RECOGNIZED_DIR.mkdir(
        parents=True,
        exist_ok=True
    )

    if not MODEL_PATH.exists():
        raise FileNotFoundError(
            f"Model YOLO tidak ditemukan: {MODEL_PATH}\n"
            "Letakkan model di lokasi tersebut atau "
            "set ECENG_MODEL_PATH."
        )


def main():
    log("==========================================")
    log("ECENG BOTTLE RECOGNITION START")
    log("==========================================")

    check_directories()

    log(f"BASE_DIR       = {BASE_DIR}")
    log(f"CAPTURE_DIR    = {CAPTURE_DIR}")
    log(f"RECOGNIZED_DIR = {RECOGNIZED_DIR}")
    log(f"MODEL_PATH     = {MODEL_PATH}")
    log(f"CONFIDENCE     = {DEFAULT_CONFIDENCE}")
    log(f"MAX_QTY        = {MAX_BOTTLE_QTY}")
    log(
        f"STRICT_TYPE    = "
        f"{STRICT_PRODUCT_TYPE}"
    )

    log("Loading YOLO model...")
    model = YOLO(str(MODEL_PATH))
    log("YOLO model loaded.")

    while True:
        conn = None

        try:
            conn = get_connection()

            project = get_active_project(conn)

            if not project:
                log(
                    "Tidak ada ec_project "
                    "dengan is_active='Y'."
                )

                conn.close()
                time.sleep(POLL_INTERVAL)
                continue

            project_code = project["project_code"]

            log(
                f"ACTIVE PROJECT: "
                f"{project_code} | "
                f"product_id={project.get('product_id')} | "
                f"product={project.get('product_nama')}"
            )

            capture = get_new_capture(
                conn,
                project_code
            )

            conn.close()
            conn = None

            if not capture:
                time.sleep(POLL_INTERVAL)
                continue

            try:
                process_capture(
                    model,
                    project,
                    capture
                )

            except Exception as exc:
                error_message = (
                    f"{type(exc).__name__}: {exc}"
                )

                log(
                    f"PROCESS ERROR id={capture['id']}: "
                    f"{error_message}"
                )

                traceback.print_exc()

                # Tandai ERROR, tetapi JANGAN hapus
                # file capture agar bisa diperiksa.
                try:
                    error_conn = get_connection()

                    update_capture_error(
                        error_conn,
                        capture["id"],
                        error_message[:255]
                    )

                    error_conn.close()

                except Exception as db_exc:
                    log(
                        "Gagal update status ERROR: "
                        f"{db_exc}"
                    )

        except KeyboardInterrupt:
            log("Program dihentikan oleh user.")
            break

        except Exception as exc:
            log(
                f"MAIN LOOP ERROR: "
                f"{type(exc).__name__}: {exc}"
            )
            traceback.print_exc()

            if conn:
                try:
                    conn.close()
                except Exception:
                    pass

            time.sleep(
                max(POLL_INTERVAL, 3)
            )


if __name__ == "__main__":
    main()
