#!/usr/bin/env python3
"""
Chuẩn hóa overlap ở ranh giới chunk (~302ms) → gap đều 30ms (khớp *_final_vi.srt).

Logic đồng nhất với fix_youtube_srt.py:fix_overlaps:
  1. đẩy start nếu start < end(prev) + gap
  2. clamp end ≤ start(kế) - gap
  3. đảm bảo tối thiểu 500ms

Sửa: *_final.srt + *_srt_map.json (Myanmar). Việt SRT đã sạch nên bỏ qua.

Usage:
    python3 normalize_overlaps.py --dry-run
    python3 normalize_overlaps.py
"""
import argparse
import json
import os
import re
import shutil
import time

PROJ = "/home/tuan-nguyen/.openclaw/workspace/015-phu_de_video"
OUT = os.path.join(PROJ, "output")
PREFIXES = ["001", "002"]
GAP_MS = 30
MIN_DUR_MS = 500

TS_RE = re.compile(r'^(\d{1,2}):(\d{2}):(\d{2}),(\d{3})\s*-->\s*(\d{1,2}):(\d{2}):(\d{2}),(\d{3})$')


def to_ms(h, m, s, ms):
    return h * 3600000 + m * 60000 + s * 1000 + ms


def fmt_ms(ms):
    ms = int(round(ms))
    h = ms // 3600000; ms %= 3600000
    m = ms // 60000; ms %= 60000
    s = ms // 1000; ms %= 1000
    return f"{h:02d}:{m:02d}:{s:02d},{ms:03d}"


def fix_overlaps(segs, gap_ms=GAP_MS):
    """segs: list of [start_ms, end_ms]. Trả về list mới, áp start-push + end-clamp + min-duration."""
    fixed = []
    for i, (start, end) in enumerate(segs):
        if i > 0 and start < fixed[-1][1] + gap_ms:
            start = fixed[-1][1] + gap_ms
        if i < len(segs) - 1 and end > segs[i + 1][0] - gap_ms:
            end = segs[i + 1][0] - gap_ms
        if end <= start:
            end = start + MIN_DUR_MS
        fixed.append([start, end])
    return fixed


def parse_srt(path):
    segs = []
    texts = []
    for blk in open(path, encoding="utf-8").read().split("\n\n"):
        lines = blk.split("\n")
        if len(lines) < 3:
            continue
        m = TS_RE.match(lines[1].strip())
        if not m:
            continue
        g = [int(x) for x in m.groups()]
        segs.append([to_ms(g[0], g[1], g[2], g[3]), to_ms(g[4], g[5], g[6], g[7])])
        texts.append("\n".join(lines[2:]).strip())
    return segs, texts


def write_srt(path, segs, texts):
    out = [f"{i}\n{fmt_ms(s[0])} --> {fmt_ms(s[1])}\n{t}"
           for i, (s, t) in enumerate(zip(segs, texts), 1)]
    with open(path, "w", encoding="utf-8") as f:
        f.write("\n\n".join(out) + "\n")


def normalize_srt_map(path):
    with open(path, encoding="utf-8") as f:
        recs = json.load(f)
    ts2ms = lambda t: to_ms(*[int(x) for x in re.split(r'[:,]', t)])
    segs = [[ts2ms(r["start"]), ts2ms(r["end"])] for r in recs]
    fixed = fix_overlaps(segs)
    for r, (st, en) in zip(recs, fixed):
        r["start"] = fmt_ms(st)
        r["end"] = fmt_ms(en)
    with open(path, "w", encoding="utf-8") as f:
        json.dump(recs, f, ensure_ascii=False, indent=2)
    return len(recs)


def backup(path):
    dst = f"{path}.bak.{time.strftime('%Y%m%d_%H%M%S')}"
    shutil.copy2(path, dst)
    return dst


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--dry-run", action="store_true")
    args = ap.parse_args()

    for p in PREFIXES:
        srt = os.path.join(OUT, f"{p}_final.srt")
        mp = os.path.join(OUT, f"{p}_srt_map.json")
        segs, texts = parse_srt(srt)
        ov_before = sum(1 for i in range(len(segs) - 1) if segs[i][1] > segs[i + 1][0])
        fixed = fix_overlaps(segs)
        ov_after = sum(1 for i in range(len(fixed) - 1) if fixed[i][1] > fixed[i + 1][0])
        if args.dry_run:
            print(f"[dry-run] {p}: {len(segs)} cues, overlaps {ov_before}→{ov_after}")
            continue
        b = backup(srt)
        write_srt(srt, fixed, texts)
        b2 = backup(mp)
        normalize_srt_map(mp)
        print(f"✅ {p}_final.srt: overlaps {ov_before}→{ov_after}  [backup {os.path.basename(b)}]")
        print(f"✅ {p}_srt_map.json updated  [backup {os.path.basename(b2)}]")


if __name__ == "__main__":
    main()
