#!/usr/bin/env python3
"""
Fix YouTube compliance cho file .srt (bản ổn định, split 2-pha cân bằng):
  1. Xóa overlap timestamp (clamp end <= start của subtitle kế tiếp)
  2. Tách subtitle quá dài: (duration > MAX_DUR) HOẶC (text > 2 dòng mà duration đủ dài)
     → split thành 2 phần tại ranh giới câu (ưu tiên gần giữa), phân bố thời gian theo tỷ lệ ký tự.
  3. Xuống dòng text > MAX_CHARS ký tự thành tối đa 2 dòng (break theo từ)

KHÔNG cắt theo dấu phẩy (tránh vỡ vụn). KHÔNG làm mất nội dung.

Usage:
    python3 fix_youtube_srt.py input.srt [-o output.srt] [--max-dur 7] [--max-chars 42]
"""
import argparse, re, math, os, shutil, time

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 parse_srt(path):
    raw = open(path, encoding='utf-8').read().rstrip()
    segs = []
    for blk in raw.split('\n\n'):
        lines = blk.split('\n')
        if len(lines) < 3:
            continue
        m = TS_RE.match(lines[1].strip())
        if not m:
            continue
        h, mm, ss, ms = int(m.group(1)), int(m.group(2)), int(m.group(3)), int(m.group(4))
        start = h*3600000 + mm*60000 + ss*1000 + ms
        h, mm, ss, ms = int(m.group(5)), int(m.group(6)), int(m.group(7)), int(m.group(8))
        end = h*3600000 + mm*60000 + ss*1000 + ms
        text = '\n'.join(lines[2:]).strip()
        if text:
            segs.append({'start': start, 'end': end, 'text': text})
    return segs

def fmt_ms(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=30):
    """Clamp overlap + đảm bảo khoảng nghỉ tối thiểu giữa 2 subtitle."""
    fixed = []
    for i, s in enumerate(segs):
        start, end = s['start'], s['end']
        # start phải >= end của segment trước + gap
        if i > 0 and start < fixed[-1]['end'] + gap_ms:
            start = fixed[-1]['end'] + gap_ms
        # end phải <= start của segment kế - gap
        if i < len(segs) - 1 and end > segs[i+1]['start'] - gap_ms:
            end = segs[i+1]['start'] - gap_ms
        # đảm bảo tối thiểu 500ms
        if end <= start:
            end = start + 500
        fixed.append({'start': start, 'end': end, 'text': s['text']})
    return fixed

def split_parts(text):
    """Tách text thành TỐI ĐA 2 phần tại ranh giới câu gần giữa (hoặc midpoint theo từ)."""
    t = text.strip()
    boundaries = [m.end() for m in re.finditer(r'(?<=[.!?;])\s+', t)]
    if boundaries:
        mid = len(t) // 2
        # chỉ chọn boundary tạo 2 phần cân đối (mỗi phần >= 25% text)
        good = [b for b in boundaries if 0.25 * len(t) <= b <= 0.75 * len(t)]
        if good:
            best = min(good, key=lambda b: abs(b - mid))
            left, right = t[:best].strip(), t[best:].strip()
            if left and right:
                return [left, right]
    words = t.split()
    if len(words) >= 4:
        mid = len(words) // 2
        return [' '.join(words[:mid]), ' '.join(words[mid:])]
    return [t]

def process_segment(start, end, text, max_dur_ms, max_chars):
    dur = end - start
    too_long = (dur > max_dur_ms) or (len(text) > max_chars * 2 and dur >= 2500)
    if not too_long:
        return [{'start': start, 'end': end, 'text': text}]
    parts = split_parts(text)
    if len(parts) < 2:
        return [{'start': start, 'end': end, 'text': text}]
    total = sum(len(p) for p in parts)
    out = []
    cur = start
    for i, p in enumerate(parts):
        if i == len(parts) - 1:
            pe = end
        else:
            pe = cur + max(1200, int(dur * len(p) / total))
        if pe > end:
            pe = end
        out.append({'start': cur, 'end': pe, 'text': p})
        cur = pe
    result = []
    for o in out:
        result.extend(process_segment(o['start'], o['end'], o['text'], max_dur_ms, max_chars))
    return result

def merge_short(segs, min_ms=800):
    """Gộp segment quá ngắn (<min_ms) vào segment trước đó."""
    out = []
    for s in segs:
        if out and (s['end'] - s['start'] < min_ms):
            prev = out[-1]
            prev['end'] = s['end']
            prev['text'] = (prev['text'] + ' ' + s['text']).strip()
        else:
            out.append(dict(s))
    return out

def wrap_text(text, max_chars):
    if len(text) <= max_chars:
        return text
    words = text.split()
    lines = []
    cur = ''
    for w in words:
        if not cur:
            cur = w
        elif len(cur) + 1 + len(w) <= max_chars:
            cur += ' ' + w
        else:
            lines.append(cur)
            cur = w
    if cur:
        lines.append(cur)
    if len(lines) > 2:
        lines = [lines[0], ' '.join(lines[1:])]
    return '\n'.join(lines)

def stats(segs, max_dur_ms, max_chars):
    ov = sum(1 for i in range(1, len(segs)) if segs[i]['start'] < segs[i-1]['end'])
    ld = sum(1 for s in segs if s['end'] - s['start'] > max_dur_ms)
    ll = sum(1 for s in segs if any(len(l) > max_chars for l in s['text'].split('\n')))
    short = sum(1 for s in segs if s['end'] - s['start'] < 800)
    return ov, ld, ll, short

def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('input')
    ap.add_argument('-o', '--output', default=None)
    ap.add_argument('--max-dur', type=float, default=7.0)
    ap.add_argument('--max-chars', type=int, default=42)
    ap.add_argument('--split', action='store_true',
                    help='(tùy chọn) tách segment dài >7s — mặc định KHÔNG tách, giữ nguyên câu dài')
    args = ap.parse_args()

    segs = parse_srt(args.input)
    max_dur_ms = int(args.max_dur * 1000)
    print(f"Đọc {len(segs)} segment từ {args.input}")
    ov, ld, ll, sh = stats(segs, max_dur_ms, args.max_chars)
    print(f"  TRƯỚC: Overlap={ov} | >{args.max_dur:.0f}s={ld} | dòng>{args.max_chars}={ll} | <0.8s={sh}")

    segs = fix_overlaps(segs)
    if args.split:
        new_segs = []
        for s in segs:
            new_segs.extend(process_segment(s['start'], s['end'], s['text'], max_dur_ms, args.max_chars))
        segs = new_segs
        for s in segs:
            s['text'] = wrap_text(s['text'], args.max_chars)
        segs = fix_overlaps(segs)
        segs = merge_short(segs)
        for s in segs:
            s['text'] = wrap_text(s['text'], args.max_chars)

    ov, ld, ll, sh = stats(segs, max_dur_ms, args.max_chars)
    print(f"  SAU  : Overlap={ov} | >{args.max_dur:.0f}s={ld} | dòng>{args.max_chars}={ll} | <0.8s={sh}")

    out = args.output or args.input
    if out == args.input:
        bak = f"{args.input}.bak.{int(time.time())}"
        shutil.copy2(args.input, bak)
        print(f"  Backup → {bak}")
    with open(out, 'w', encoding='utf-8') as f:
        for i, s in enumerate(segs, 1):
            f.write(f"{i}\n{fmt_ms(s['start'])} --> {fmt_ms(s['end'])}\n{s['text']}\n\n")
    print(f"→ {out} ({len(segs)} segment)")

if __name__ == '__main__':
    main()
