#!/usr/bin/env python3
"""Bulk dump ALL tables from yajny + analytics databases via docker exec artisan tinker.
Output: CSV files in db_dump/{db}/{table}.csv + schema files + manifest.
Skips tables already dumped (resumable).
"""
import subprocess, base64, sys, json, csv, os, time

HOST = "root@46.101.74.150"
KEY = "/root/ir-assessment/redteam/yajny/L3/ssh_keys/jenkins-host-id_rsa.pem"
CONTAINER = "yajny-api"
DUMP_DIR = "/root/ir-assessment/redteam/yajny/db_dump"

def run_tinker(queries, timeout=600):
    code = "\n".join(queries)
    b64 = base64.b64encode(code.encode()).decode()
    cmd = [
        "ssh", "-i", KEY, "-o", "BatchMode=yes", "-o", "StrictHostKeyChecking=accept-new",
        HOST,
        f'echo {b64} | base64 -d | docker exec -i {CONTAINER} php /var/www/html/artisan tinker 2>&1'
    ]
    r = subprocess.run(cmd, capture_output=True, text=True, timeout=timeout)
    return r.stdout

def get_table_list():
    """Get all tables with row counts from information_schema."""
    queries = [
        "$ts = DB::select('SELECT TABLE_NAME, TABLE_ROWS, DATA_LENGTH FROM information_schema.TABLES WHERE TABLE_SCHEMA = \"yajny\" ORDER BY TABLE_NAME');",
        'foreach($ts as $t) { echo $t->TABLE_NAME."|".$t->TABLE_ROWS."|".$t->DATA_LENGTH."\n"; }',
        'exit;',
    ]
    raw = run_tinker(queries, timeout=120)
    tables = []
    for line in raw.split("\n"):
        line = line.strip()
        if "|" in line and not line.startswith(">") and not line.startswith("echo") and "TABLE_NAME" not in line:
            parts = line.split("|")
            if len(parts) == 3 and parts[0].replace("_","").isalnum():
                tables.append({
                    "name": parts[0],
                    "rows": parts[1] if parts[1] != "" else "0",
                    "data_length": parts[2] if parts[2] != "" else "0",
                })
    return tables

def get_schema(table):
    """Get DESCRIBE for a table."""
    queries = [
        f'$c = DB::select("DESCRIBE `{table}`"); foreach($c as $col) {{ echo $col->Field."|".$col->Type."|".$col->Null."|".$col->Key."|".$col->Default."|".$col->Extra."\n"; }}',
        'exit;',
    ]
    return run_tinker(queries, timeout=60)

def dump_table(table, max_rows=50000):
    """Dump table to CSV via paginated SELECT *."""
    csv_path = f"{DUMP_DIR}/yajny/{table}.csv"
    if os.path.isfile(csv_path) and os.path.getsize(csv_path) > 0:
        return "SKIP"
    
    # Skip rep_ tables (VIEWs that cause hangs)
    if table.startswith("rep_"):
        return "SKIP_VIEW"
    
    # First get columns
    schema_raw = get_schema(table)
    columns = []
    for line in schema_raw.split("\n"):
        line = line.strip()
        if "|" in line and not line.startswith(">") and not line.startswith("echo") and "Field" not in line:
            parts = line.split("|")
            if len(parts) >= 2 and parts[0].replace("_","").isalnum():
                columns.append(parts[0])
    
    if not columns:
        return "NO_SCHEMA"
    
    # Save schema
    with open(f"{DUMP_DIR}/yajny/{table}_schema.txt", "w") as f:
        f.write(schema_raw)
    
    # Paginated dump using keyset pagination (WHERE id > last_id)
    # Faster than OFFSET on large tables
    total_rows = 0
    write_header = True
    last_id = None
    id_col = "id" if "id" in columns else columns[0]
    
    while True:
        if last_id is not None:
            where_clause = f"WHERE `{id_col}` > {last_id}"
        else:
            where_clause = ""
        col_list = ",".join(f"`{c}`" for c in columns)
        queries = [
            f'$r = DB::select("SELECT {col_list} FROM `{table}` {where_clause} ORDER BY `{id_col}` LIMIT 500");',
            'foreach($r as $row) { echo json_encode($row)."\n"; }',
        ]
        raw = run_tinker(queries, timeout=300)
        
        # Parse JSON lines
        page_rows = 0
        with open(csv_path, "a", newline="") as f:
            writer = csv.writer(f)
            if write_header:
                writer.writerow(columns)
                write_header = False
            for line in raw.split("\n"):
                line = line.strip()
                if len(line) < 5 or not line.startswith("{"):
                    continue
                # Skip Psy Shell object dumps ({#1025)
                if line.startswith("{#"):
                    continue
                try:
                    obj = json.loads(line)
                    if not isinstance(obj, dict):
                        continue
                    row = []
                    for c in columns:
                        val = obj.get(c, "")
                        if isinstance(val, (dict, list)):
                            val = json.dumps(val, ensure_ascii=False)
                        row.append(str(val) if val is not None else "")
                    writer.writerow(row)
                    page_rows += 1
                except:
                    pass
        
        total_rows += page_rows
        if page_rows < 500:
            break
        
        # Update last_id from the last row in this batch
        # Parse it from the raw output directly
        last_line = None
        for line in raw.split("\n"):
            line = line.strip()
            if len(line) < 5 or not line.startswith("{"):
                continue
            if line.startswith("{#"):
                continue
            try:
                obj = json.loads(line)
                if isinstance(obj, dict):
                    last_line = obj
            except:
                pass
        
        if last_line and id_col in last_line:
            val = last_line[id_col]
            if isinstance(val, (int, str)) and str(val).isdigit():
                last_id = int(val)
            else:
                break
        else:
            break
        
        # Safety limit
        if total_rows >= max_rows:
            break
        
        # Progress
        if total_rows % 5000 == 0 or total_rows % 50000 == 0:
            print(f"    {table}: {total_rows} rows...")
    
    return total_rows

if __name__ == "__main__":
    print("=" * 60)
    print("BULK DUMP — ALL TABLES")
    print("=" * 60)
    
    tables = get_table_list()
    print(f"[*] Found {len(tables)} tables")
    
    manifest = []
    for i, t in enumerate(tables):
        name = t["name"]
        rows = t["rows"]
        size = int(t["data_length"]) if t["data_length"].isdigit() else 0
        
        # Skip very large tables (>50GB) — won't fit
        if size > 50_000_000_000:
            print(f"[{i+1}/{len(tables)}] SKIP {name} ({rows} rows, {size//1024//1024//1024}GB — too large)")
            manifest.append({"table": name, "rows": rows, "size_bytes": size, "status": "SKIP_TOO_LARGE"})
            continue
        
        # Cap at 50K rows for medium tables
        max_r = 50000 if size > 500_000_000 else 1000000
        
        print(f"[{i+1}/{len(tables)}] {name} ({rows} rows, {size//1024}KB)")
        result = dump_table(name, max_rows=max_r)
        actual = result if isinstance(result, int) else 0
        status = "OK" if isinstance(result, int) else result
        print(f"    → {actual} rows dumped ({status})")
        manifest.append({"table": name, "rows": rows, "size_bytes": size, "dumped": actual, "status": status})
    
    # Save manifest
    with open(f"{DUMP_DIR}/yajny/_manifest.json", "w") as f:
        json.dump(manifest, f, indent=2)
    
    # Summary
    total_dumped = sum(m.get("dumped",0) for m in manifest)
    total_skipped = sum(1 for m in manifest if m["status"] != "OK")
    print(f"\n{'='*60}")
    print(f"COMPLETE: {total_dumped} rows dumped, {total_skipped} tables skipped")
    print(f"Manifest: {DUMP_DIR}/yajny/_manifest.json")
    print(f"{'='*60}")
