#!/usr/bin/env python3
"""
Stream-analyze a MySQL dump for secrets/PII.
For each dump:
  1. List all tables (CREATE TABLE)
  2. Find CREATE TABLE definitions containing secret columns
  3. Extract INSERT INTO data for tables with secret columns
  4. Count total rows per table
  5. Report findings with exact values (NO REDACTION)
"""
import sys
import re
import os
import json
from collections import defaultdict

SECRET_PATTERNS = [
    'password', 'passwd', 'pwd', 'user_pass',
    'secret', 'token', 'api_key', 'apikey', 'access_token',
    'refresh_token', 'credential', 'private_key',
    'session', 'cookie', 'hash', 'salt'
]

# Words to exclude from triggering (avoid false positives on column names like 'tokenization')
# We'll match tokens as substrings of column names but exclude pure word "tokenize"
EXCLUDE_WORDS = {'tokenize', 'tokenized', 'tokenization', 'hashed', 'hashtag'}

def is_secret_column(col_name_lower):
    """Check if a column name indicates a secret-bearing column."""
    for pat in SECRET_PATTERNS:
        if pat in col_name_lower:
            # exclude false positives
            if any(ex in col_name_lower for ex in EXCLUDE_WORDS):
                continue
            return True
    return False

CREATE_RE = re.compile(r'^\s*CREATE\s+TABLE\s+(?:IF\s+NOT\s+EXISTS\s+)?`?(\w+)`?\s*\(', re.IGNORECASE)
INSERT_RE = re.compile(r'^\s*INSERT\s+INTO\s+`?(\w+)`?\s+', re.IGNORECASE)

def analyze_dump(path, out_dir):
    basename = os.path.basename(path)
    out_path = os.path.join(out_dir, basename + '.report.md')
    
    tables = []                       # list of table names in order
    table_secret_cols = {}            # table_name -> list of secret-bearing column names
    table_row_counts = defaultdict(int)
    secret_inserts = defaultdict(list)  # table_name -> list of INSERT lines (first 20)
    secret_insert_count = defaultdict(int)  # table_name -> total count of INSERT lines with secrets
    
    # State for parsing CREATE TABLE bodies
    in_create = False
    current_table = None
    create_body = []
    
    # Regex for column definitions inside CREATE TABLE
    col_def_re = re.compile(r'^\s*`?(\w+)`?\s+(?:tinyint|smallint|mediumint|int|bigint|varchar|char|text|tinytext|mediumtext|longtext|blob|tinyblob|mediumblob|longblob|date|datetime|timestamp|time|year|decimal|float|double|enum|set|json|bool|boolean)', re.IGNORECASE)
    
    print(f"[*] Analyzing {basename} ...", flush=True)
    
    with open(path, 'r', errors='replace') as f:
        for line in f:
            # CREATE TABLE detection
            m = CREATE_RE.match(line)
            if m:
                current_table = m.group(1)
                tables.append(current_table)
                in_create = True
                create_body = []
                continue
            if in_create:
                # End of CREATE TABLE (ends with ; on its own line typically)
                if re.match(r'^\s*\)\s*ENGINE', line, re.IGNORECASE) or re.match(r'^\s*\)\s*;', line) or (line.strip().endswith(';') and '`' not in line and '(' not in line and 'INSERT' not in line.upper()):
                    # parse body for secret columns
                    cols = []
                    for bline in create_body:
                        cm = col_def_re.match(bline)
                        if cm:
                            cols.append(cm.group(1))
                    secret_cols = [c for c in cols if is_secret_column(c.lower())]
                    if secret_cols:
                        table_secret_cols[current_table] = secret_cols
                    in_create = False
                    current_table = None
                    create_body = []
                    continue
                create_body.append(line)
                # also catch single-line bodies
                continue
            
            # INSERT INTO detection
            im = INSERT_RE.match(line)
            if im:
                tname = im.group(1)
                table_row_counts[tname] += 1
                # if this table has secret columns, capture the INSERT line
                if tname in table_secret_cols:
                    secret_insert_count[tname] += 1
                    if len(secret_inserts[tname]) < 20:
                        secret_inserts[tname].append(line.rstrip('\n'))
    
    # Write report
    with open(out_path, 'w') as out:
        out.write(f"# Analysis: {basename}\n\n")
        out.write(f"**File size:** {os.path.getsize(path):,} bytes\n")
        out.write(f"**Total tables:** {len(tables)}\n")
        out.write(f"**Tables with secret columns:** {len(table_secret_cols)}\n")
        total_inserts = sum(table_row_counts.values())
        out.write(f"**Total INSERT statements:** {total_inserts:,}\n\n")
        
        out.write("## All Tables (with row counts)\n\n")
        out.write("| # | Table | INSERT count | Secret columns |\n")
        out.write("|---|-------|--------------|----------------|\n")
        for i, t in enumerate(tables, 1):
            cnt = table_row_counts.get(t, 0)
            sc = ', '.join(table_secret_cols.get(t, [])) or '-'
            out.write(f"| {i} | `{t}` | {cnt:,} | {sc} |\n")
        
        out.write("\n## Secret Column Details\n\n")
        if not table_secret_cols:
            out.write("No tables with secret-bearing columns detected.\n")
        else:
            for t, cols in table_secret_cols.items():
                out.write(f"### `{t}`\n")
                out.write(f"- Secret columns: {', '.join(cols)}\n")
                out.write(f"- INSERT statements: {secret_insert_count.get(t, 0):,}\n")
                if secret_inserts.get(t):
                    out.write(f"- Sample INSERT data (first {min(20, len(secret_inserts[t]))} rows):\n\n")
                    out.write("```\n")
                    for s in secret_inserts[t]:
                        out.write(s + "\n")
                    out.write("```\n")
                out.write("\n")
    
    # Return summary for parent script
    summary = {
        'file': basename,
        'size': os.path.getsize(path),
        'tables': len(tables),
        'tables_with_secrets': len(table_secret_cols),
        'secret_tables': list(table_secret_cols.keys()),
        'secret_columns': {t: table_secret_cols[t] for t in table_secret_cols},
        'secret_insert_count': dict(secret_insert_count),
        'row_counts': dict(table_row_counts),
        'report_path': out_path,
    }
    return summary

if __name__ == '__main__':
    dumps = sys.argv[1:]
    out_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)))
    os.makedirs(out_dir, exist_ok=True)
    all_summaries = []
    for d in dumps:
        if not os.path.isfile(d):
            print(f"[!] Skip (not found): {d}", flush=True)
            continue
        try:
            s = analyze_dump(d, out_dir)
            all_summaries.append(s)
            print(f"[+] {s['file']}: {s['tables']} tables, {s['tables_with_secrets']} with secrets", flush=True)
        except Exception as e:
            print(f"[!] Error on {d}: {e}", flush=True)
    # write consolidated JSON
    with open(os.path.join(out_dir, '_summary.json'), 'w') as f:
        json.dump(all_summaries, f, indent=2)
    print(f"[+] Done. {len(all_summaries)} dumps analyzed.", flush=True)
