import re
import mysql.connector
from collections import defaultdict
from datetime import datetime, timedelta, date as date_cls

CONFIG_PATH = "/var/www/html/projects/aerion/global_dbconfig.php"

def extract_config(path):
    with open(path, "r") as f:
        content = f.read()

    def get_value(key):
        pattern = rf"\$GLOBALS\['{key}'\]\s*=\s*\"(.*?)\";"
        match = re.search(pattern, content)
        return match.group(1) if match else None

    return {
        "host": get_value("ser_sh"),
        "user": get_value("UN"),
        "password": get_value("pswd"),
        "database": get_value("DB"),
    }

def parse_date(d):
    if not d:
        return None
    if isinstance(d, date_cls):
        return d
    if isinstance(d, datetime):
        return d.date()
    try:
        return datetime.strptime(str(d), "%Y-%m-%d").date()
    except:
        return None

read_cfg = extract_config(CONFIG_PATH)

write_cfg = {
    "host": "dev1.atguru.work",
    "user": "hemanta",
    "password": "Hemnt@01",
    "database": "hemanta"
}

conn_read = mysql.connector.connect(**read_cfg)
cursor_read = conn_read.cursor()

conn_write = mysql.connector.connect(**write_cfg)
cursor_write = conn_write.cursor()

try:
    cursor_read.execute("""
        SELECT net_attr_st_dt
        FROM aerion.net_data
        WHERE net_script_nm = 'dev_attr_est_reqs_stat_forcast'
        LIMIT 1
    """)

    row = cursor_read.fetchone()
    if not row or not row[0]:
        print("[ERROR] Missing affected date")
        raise SystemExit(1)

    AFFECTED_START_DATE = parse_date(row[0])
    if not AFFECTED_START_DATE:
        print("[ERROR] Invalid affected date")
        raise SystemExit(1)

    print(f"[INFO] Start date: {AFFECTED_START_DATE}")

    cursor_read.execute("""
        SELECT for_stat_st_dt_fmt2, for_stat_e_dt_fmt2
        FROM aerion.vw_dev_est_reqs_status_forcast
    """)

    rows = cursor_read.fetchall()

    new_count = defaultdict(int)
    close_count = defaultdict(int)
    source_dates = set()
    null_start_count = 0
    null_end_count = 0

    for start, end in rows:
        s = parse_date(start)
        e = parse_date(end)

        if s:
            new_count[s] += 1
            source_dates.add(s)
        else:
            null_start_count += 1

        if e:
            close_count[e] += 1
            source_dates.add(e)
        else:
            null_end_count += 1

    if not source_dates:
        print("[ERROR] No valid source data")
        raise SystemExit(1)

    max_date = max(source_dates)
    next_date = max_date + timedelta(days=1)

    print(f"[INFO] Source max date: {max_date}")

    cursor_write.execute("""
        SELECT iddev_attr_est_reqs_stat_forcast, dev_est_reqs_stat_forcast_dt
        FROM dev_attr_est_reqs_stat_forcast
        WHERE dev_est_reqs_stat_forcast_dt >= %s
    """, (AFFECTED_START_DATE,))

    existing_map = {}
    for rid, d in cursor_write.fetchall():
        pd = parse_date(d)
        if pd:
            existing_map[pd] = rid

    cursor_write.execute("""
        SELECT total
        FROM dev_attr_est_reqs_stat_forcast
        WHERE qry_dt < %s
          AND category = 'Total'
        ORDER BY qry_dt DESC
        LIMIT 1
    """, (AFFECTED_START_DATE,))

    row = cursor_write.fetchone()
    running_total = row[0] if row else 0

    cursor_write.execute("""
        UPDATE dev_attr_est_reqs_stat_forcast
        SET qry_dt = NULL,
            `new` = NULL,
            `close` = NULL,
            total = NULL,
            new_pcnt = NULL,
            close_pcnt = NULL,
            change_pcnt = NULL,
            category = NULL,
            section1 = NULL,
            idparent = NULL
        WHERE dev_est_reqs_stat_forcast_dt >= %s
    """, (AFFECTED_START_DATE,))

    print(f"[INFO] Rows reset: {cursor_write.rowcount}")

    all_dates = sorted(source_dates | set(existing_map.keys()) | {AFFECTED_START_DATE, next_date})

    update_rows = []
    insert_rows = []
    nullified_rows = 0

    for d in all_dates:
        if d < AFFECTED_START_DATE:
            continue

        if d == next_date:
            new = null_start_count
            close = null_end_count
            qry_dt = None
            running_total += new - close
            row_data = (
                d,
                qry_dt,
                new,
                close,
                running_total,
                f"{new * 100:.2f}%",
                f"{close * 100:.2f}%",
                f"{(new - close) * 100:.2f}%",
                "Total",
                "Total",
                None
            )
        elif d not in source_dates:
            nullified_rows += 1
            row_data = (
                d,
                None,
                None,
                None,
                None,
                None,
                None,
                None,
                None,
                None,
                None
            )
        else:
            new = new_count[d]
            close = close_count[d]
            qry_dt = d
            running_total += new - close
            row_data = (
                d,
                qry_dt,
                new,
                close,
                running_total,
                f"{new * 100:.2f}%",
                f"{close * 100:.2f}%",
                f"{(new - close) * 100:.2f}%",
                "Total",
                "Total",
                None
            )

        existing_id = existing_map.get(d)

        if existing_id is not None:
            update_rows.append(row_data + (existing_id,))
        else:
            insert_rows.append(row_data)

    print(f"[INFO] Rows to update: {len(update_rows)}")
    print(f"[INFO] Rows to insert: {len(insert_rows)}")

    if update_rows:
        cursor_write.executemany("""
            UPDATE dev_attr_est_reqs_stat_forcast
            SET qry_dt = %s,
                `new` = %s,
                `close` = %s,
                total = %s,
                new_pcnt = %s,
                close_pcnt = %s,
                change_pcnt = %s,
                category = %s,
                section1 = %s,
                idparent = %s
            WHERE iddev_attr_est_reqs_stat_forcast = %s
        """, [
            (r[1], r[2], r[3], r[4], r[5], r[6], r[7], r[8], r[9], r[10], r[11])
            for r in update_rows
        ])

    if insert_rows:
        cursor_write.executemany("""
            INSERT INTO dev_attr_est_reqs_stat_forcast (
                dev_est_reqs_stat_forcast_dt,
                qry_dt,
                `new`,
                `close`,
                total,
                new_pcnt,
                close_pcnt,
                change_pcnt,
                category,
                section1,
                idparent
            ) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
        """, insert_rows)

    conn_write.commit()

    checkpoint_value = max_date.strftime("%Y-%m-%d")
    cursor_read.execute("""
        UPDATE aerion.net_data
        SET net_attr_st_dt = %s
        WHERE net_script_nm = 'dev_attr_est_reqs_stat_forcast'
    """, (checkpoint_value,))

    conn_read.commit()

    cursor_read.execute("""
        SELECT net_attr_st_dt
        FROM aerion.net_data
        WHERE net_script_nm = 'dev_attr_est_reqs_stat_forcast'
        LIMIT 1
    """)

    updated_date = cursor_read.fetchone()[0]
    print(f"[INFO] Updated: {len(update_rows)} | Inserted: {len(insert_rows)} | Checkpoint: {updated_date}")

except Exception as e:
    try:
        conn_write.rollback()
    except:
        pass
    try:
        conn_read.rollback()
    except:
        pass
    print(f"[ERROR] {e}")
    raise SystemExit(1)

finally:
    cursor_read.close()
    conn_read.close()
    cursor_write.close()
    conn_write.close()
