#!/usr/bin/env python3
"""Stream-extract INSERT INTO statements for specific tables in large MySQL dumps.
Handles multi-line INSERTs by buffering until the statement terminator ';'.
Outputs first N INSERT statements per table with no redaction."""
import sys, re, os

def extract_inserts(path, tables, max_per_table=5, max_chars_per_stmt=8000):
    """Stream through file, capture INSERT INTO `table` ... ; statements."""
    wanted = {t.lower() for t in tables}
    results = {t: [] for t in tables}
    counts = {t: 0 for t in tables}
    
    # Build regex for matching INSERT INTO `table`
    insert_re = re.compile(r'^\s*INSERT\s+INTO\s+`?(\w+)`?\s+', re.IGNORECASE)
    
    buf = None
    buf_table = None
    with open(path, 'r', errors='replace') as f:
        for line in f:
            if buf is not None:
                buf += line
                # Check if statement terminated (line ends with ; and not inside string)
                # Heuristic: line ends with ; followed by newline
                stripped = line.rstrip('\n').rstrip('\r')
                if stripped.endswith(';'):
                    # finalize
                    counts[buf_table] = counts.get(buf_table, 0) + 1
                    if len(results.get(buf_table, [])) < max_per_table:
                        disp = buf if len(buf) < max_chars_per_stmt else buf[:max_chars_per_stmt] + " ...[truncated]"
                        results.setdefault(buf_table, []).append(disp)
                    buf = None
                    buf_table = None
                continue
            m = insert_re.match(line)
            if m:
                tname = m.group(1)
                if tname.lower() in wanted:
                    buf = line
                    buf_table = tname
                    # check single-line termination
                    stripped = line.rstrip('\n').rstrip('\r')
                    if stripped.endswith(';'):
                        counts[buf_table] = counts.get(buf_table, 0) + 1
                        if len(results.get(buf_table, [])) < max_per_table:
                            disp = buf if len(buf) < max_chars_per_stmt else buf[:max_chars_per_stmt] + " ...[truncated]"
                            results.setdefault(buf_table, []).append(disp)
                        buf = None
                        buf_table = None
    return results, counts

if __name__ == '__main__':
    # Usage: extract_inserts.py <dump> <table1> [table2 ...]
    path = sys.argv[1]
    tables = sys.argv[2:]
    results, counts = extract_inserts(path, tables)
    for t in tables:
        print(f"\n=== {t} ({counts.get(t,0)} total INSERTs) ===")
        for r in results.get(t, []):
            print(r)
