#!/usr/bin/env python3
"""
tg_txt_decrypt.py — Telegram DNS TXT config decryptor

Fetches and decrypts help.configSimple from Telegram's bootstrap DNS TXT records:
  apv3.stel.com  (production, current)
  apv2.stel.com  (production, legacy)
  tapv3.stel.com (test)
  tapv2.stel.com (test, legacy)

No DoH required. Uses raw UDP/TCP DNS via dnspython, or falls back to
passing TXT values directly on the command line.

Dependencies:
  pip install dnspython pycryptodome cryptography

Usage:
  python3 tg_txt_decrypt.py                          # query all default domains
  python3 tg_txt_decrypt.py apv3.stel.com            # query specific domain
  python3 tg_txt_decrypt.py --nameserver 1.1.1.1 apv3.stel.com
  python3 tg_txt_decrypt.py --records "base64val1" "base64val2"  # skip DNS

References:
  tdesktop: mtproto/special_config_request.cpp  decryptSimpleConfig()
  tdesktop: mtproto/details/mtproto_tls_socket.cpp
"""

import re
import sys
import base64
import hashlib
import struct
import subprocess
import tempfile
import os
import time
import argparse
from cryptography.hazmat.primitives.serialization import load_pem_public_key

try:
    from Crypto.Cipher import AES
except ImportError:
    sys.exit("Missing pycryptodome. Run: pip install pycryptodome")

try:
    import dns.resolver
    HAS_DNSPYTHON = True
except ImportError:
    HAS_DNSPYTHON = False

# ── RSA public key ────────────────────────────────────────────────────────────
# From tdesktop: mtproto/special_config_request.cpp  (kPublicKey)
# Used for both production and test TXT records (same key, different domains).
_SPECIAL_CONFIG_RSA_PKCS1 = """\
-----BEGIN RSA PUBLIC KEY-----
MIIBCgKCAQEAyr+18Rex2ohtVy8sroGPBwXD3DOoKCSpjDqYoXgCqB7ioln4eDCF
fOBUlfXUEvM/fnKCpF46VkAftlb4VuPDeQSS/ZxZYEGqHaywlroVnXHIjgqoxiAd
192xRGreuXIaUKmkwlM9JID9WS2jUsTpzQ91L8MEPLJ/4zrBwZua8W5fECwCCh2c
9G5IzzBm+otMS/YKwmR1olzRCyEkyAEjXWqBI9Ftv5eG8m0VkBzOG655WIYdyV0H
fDK/NWcvGqa0w/nriMD6mDjKOryamw0OP9QuYgMN0C9xMW9y8SmP4h92OAWodTYg
Y1hZCxdv6cs5UnW9+PWvS+WIbkh+GaWYxwIDAQAB
-----END RSA PUBLIC KEY-----"""

DEFAULT_DOMAINS = [
    "apv3.stel.com",
    "apv2.stel.com",
    "tapv3.stel.com",
    "tapv2.stel.com",
]

# ── TL constructors ───────────────────────────────────────────────────────────
_HELP_CONFIG_SIMPLE = 0x5A592A6C
_TL_VECTOR          = 0x1CB5C415
_ACCESS_POINT_RULE  = 0x4679B65F
_IP_PORT            = 0xD433AD73
_IP_PORT_SECRET     = 0x37982646


# ── RSA ───────────────────────────────────────────────────────────────────────

def _load_public_key(pkcs1_pem: str):
    """Convert PKCS#1 PEM → cryptography public key (requires openssl CLI)."""
    with tempfile.NamedTemporaryFile(suffix='.pem', delete=False, mode='w') as f:
        f.write(pkcs1_pem)
        path = f.name
    try:
        r = subprocess.run(
            ['openssl', 'rsa', '-RSAPublicKey_in', '-in', path, '-pubout'],
            capture_output=True, check=True,
        )
        return load_pem_public_key(r.stdout)
    finally:
        os.unlink(path)


def _rsa_public_raw(data: bytes, pub) -> bytes:
    """Raw RSA public operation: c^e mod n → 256-byte result."""
    nums = pub.public_numbers()
    m = int.from_bytes(data, 'big')
    if m >= nums.n:
        raise ValueError("Ciphertext ≥ modulus — wrong data or wrong key")
    return pow(m, nums.e, nums.n).to_bytes(256, 'big')


# ── DNS ───────────────────────────────────────────────────────────────────────

def query_txt_dnspython(domain: str, nameserver: str | None = None) -> list[str]:
    if not HAS_DNSPYTHON:
        raise RuntimeError("dnspython not installed (pip install dnspython)")
    resolver = dns.resolver.Resolver(configure=False)
    resolver.nameservers = [nameserver or '8.8.8.8']
    resolver.lifetime = 10
    answer = resolver.resolve(domain, 'TXT')
    records = []
    for rdata in answer:
        for string in rdata.strings:
            records.append(string.decode('ascii', errors='replace'))
    return records


def query_txt_dig(domain: str, nameserver: str | None = None) -> list[str]:
    cmd = ['dig', '+short', domain, 'TXT']
    if nameserver:
        cmd += [f'@{nameserver}']
    r = subprocess.run(cmd, capture_output=True, text=True, check=True)
    records = []
    for line in r.stdout.splitlines():
        line = line.strip().strip('"')
        if line:
            records.append(line)
    return records


def query_txt(domain: str, nameserver: str | None = None) -> list[str]:
    """Try dnspython first, then dig."""
    if HAS_DNSPYTHON:
        return query_txt_dnspython(domain, nameserver)
    return query_txt_dig(domain, nameserver)


# ── Decryption pipeline ───────────────────────────────────────────────────────

def _concatenate(records: list[str]) -> bytes:
    """
    Replicate ConcatenateDnsTxtFields():
      sort by descending string length → join → base64-filter → expect 344 chars → decode.
    """
    joined = ''.join(sorted(records, key=lambda s: -len(s)))
    clean  = re.sub(r'[^A-Za-z0-9+/=]', '', joined)
    if len(clean) != 344:
        raise ValueError(
            f"Expected 344 base64 chars after filtering, got {len(clean)}\n"
            f"  (check that you have both TXT records)"
        )
    return base64.b64decode(clean + '==')


def _aes_cbc_decrypt_and_verify(rsa_plain: bytes) -> bytes:
    """
    Replicate decryptSimpleConfig() AES step:
      key  = rsa_plain[0:32]
      iv   = rsa_plain[16:32]   ← overlaps with key
      data = AES-CBC-decrypt(rsa_plain[32:])  → 224 bytes
      check SHA-256(data[:208])[:16] == data[208:224]
      return data[4 : 4 + data[0:4] as int32]  (TL payload)
    """
    key     = bytes(rsa_plain[0:32])
    iv      = bytes(rsa_plain[16:32])
    aes_out = AES.new(key, AES.MODE_CBC, iv).decrypt(bytes(rsa_plain[32:]))

    body           = aes_out[:208]
    digest_stored  = aes_out[208:224]
    digest_computed = hashlib.sha256(body).digest()[:16]
    if digest_stored != digest_computed:
        raise ValueError(
            f"AES digest mismatch\n"
            f"  stored  : {digest_stored.hex()}\n"
            f"  computed: {digest_computed.hex()}"
        )

    real_len = struct.unpack_from('<i', body, 0)[0]
    if not (0 < real_len <= 204) or real_len & 3:
        raise ValueError(f"Invalid TL payload length: {real_len}")
    return body[4:4 + real_len]


def decrypt_txt_records(records: list[str], pub) -> dict:
    """Full pipeline: TXT records → parsed config dict."""
    cipher_bytes = _concatenate(records)
    rsa_plain    = _rsa_public_raw(cipher_bytes, pub)
    tl_payload   = _aes_cbc_decrypt_and_verify(rsa_plain)
    return _parse_config_simple(tl_payload)


# ── TL parser ─────────────────────────────────────────────────────────────────

def _read_string(buf: bytes, off: int) -> tuple[str, int]:
    slen = buf[off]; off += 1
    s    = buf[off:off + slen].decode('utf-8', errors='replace')
    off += slen + (4 - (1 + slen) % 4) % 4
    return s, off


def _parse_config_simple(tl: bytes) -> dict:
    off = 0

    cid = struct.unpack_from('<I', tl, off)[0]; off += 4
    if cid != _HELP_CONFIG_SIMPLE:
        raise ValueError(f"Bad constructor {cid:#010x}, expected {_HELP_CONFIG_SIMPLE:#010x}")

    date_v    = struct.unpack_from('<i', tl, off)[0]; off += 4
    expires_v = struct.unpack_from('<i', tl, off)[0]; off += 4

    if struct.unpack_from('<I', tl, off)[0] == _TL_VECTOR:
        off += 4
    rule_count = struct.unpack_from('<I', tl, off)[0]; off += 4

    rules = []
    for _ in range(rule_count):
        if struct.unpack_from('<I', tl, off)[0] != _ACCESS_POINT_RULE:
            raise ValueError("Expected accessPointRule constructor")
        off += 4

        phone_prefix, off = _read_string(tl, off)
        dc_id = struct.unpack_from('<i', tl, off)[0]; off += 4

        if struct.unpack_from('<I', tl, off)[0] == _TL_VECTOR:
            off += 4
        ep_count = struct.unpack_from('<I', tl, off)[0]; off += 4

        endpoints = []
        for _ in range(ep_count):
            ip_cid = struct.unpack_from('<I', tl, off)[0]; off += 4
            ipv4   = struct.unpack_from('<I', tl, off)[0]; off += 4
            port   = struct.unpack_from('<i', tl, off)[0]; off += 4
            ip_str = '%d.%d.%d.%d' % (
                (ipv4 >> 24) & 0xFF, (ipv4 >> 16) & 0xFF,
                (ipv4 >>  8) & 0xFF,  ipv4         & 0xFF,
            )
            secret = None
            if ip_cid == _IP_PORT_SECRET:
                b0 = tl[off]; off += 1
                secret = bytes(tl[off:off + b0]); off += b0
                off += (4 - (1 + b0) % 4) % 4
            endpoints.append({'ip': ip_str, 'port': port, 'secret': secret})

        rules.append({'dc_id': dc_id, 'phone_prefix': phone_prefix, 'endpoints': endpoints})

    return {'date': date_v, 'expires': expires_v, 'rules': rules}


# ── Output ────────────────────────────────────────────────────────────────────

def _fmt_secret(secret: bytes | None) -> str:
    if secret is None:
        return '(none — plain TCP)'
    if secret[0] == 0xEE and len(secret) >= 18:
        host = secret[17:].decode('utf-8', errors='replace')
        key  = secret[1:17].hex()
        return f'FakeTLS  host={host}  key={key}'
    return secret.hex()


def print_result(domain: str, cfg: dict):
    now     = int(time.time())
    expired = now > cfg['expires']

    print(f"\n{'─'*60}")
    print(f"  Domain  : {domain}")
    print(f"  Date    : {time.strftime('%Y-%m-%d %H:%M:%S UTC', time.gmtime(cfg['date']))}")
    print(f"  Expires : {time.strftime('%Y-%m-%d %H:%M:%S UTC', time.gmtime(cfg['expires']))} "
          f"{'[EXPIRED]' if expired else '[valid]'}")

    for rule in cfg['rules']:
        prefix = repr(rule['phone_prefix']) if rule['phone_prefix'] else '(all)'
        print(f"\n  DC {rule['dc_id']}  phone_prefix={prefix}")
        for ep in rule['endpoints']:
            print(f"    address : {ep['ip']}:{ep['port']}")
            print(f"    secret  : {_fmt_secret(ep['secret'])}")
            if ep['secret']:
                print(f"    tg link : tg://proxy?server={ep['ip']}"
                      f"&port={ep['port']}&secret={ep['secret'].hex()}")


# ── CLI ───────────────────────────────────────────────────────────────────────

def build_parser() -> argparse.ArgumentParser:
    p = argparse.ArgumentParser(
        description=__doc__,
        formatter_class=argparse.RawDescriptionHelpFormatter,
    )
    p.add_argument('domains', nargs='*', default=DEFAULT_DOMAINS,
                   help='Domain(s) to query (default: all four stel.com domains)')
    p.add_argument('--nameserver', '-n', default=None,
                   help='DNS server to use (default: 8.8.8.8)')
    p.add_argument('--records', '-r', nargs='+', metavar='BASE64',
                   help='Supply raw TXT record base64 values directly, '
                        'skipping DNS. Use together with a single positional '
                        'domain for labelling, e.g.: '
                        'tg_txt_decrypt.py apv3.stel.com --records "aaa" "bbb"')
    return p


def main():
    args = build_parser().parse_args()
    pub  = _load_public_key(_SPECIAL_CONFIG_RSA_PKCS1)

    if args.records:
        # Domain positional is optional label when supplying records manually
        domain = args.domains[0] if args.domains != DEFAULT_DOMAINS else '(manual)'
        print(f"Using {len(args.records)} manually supplied record(s) for {domain}")
        try:
            cfg = decrypt_txt_records(args.records, pub)
            print_result(domain, cfg)
        except Exception as exc:
            print(f"  ERROR: {exc}", file=sys.stderr)
        return

    for domain in args.domains:
        print(f"Querying {domain} ...", end=' ', flush=True)
        try:
            records = query_txt(domain, args.nameserver)
            if not records:
                print("no TXT records returned")
                continue
            print(f"{len(records)} record(s) received")
            cfg = decrypt_txt_records(records, pub)
            print_result(domain, cfg)
        except Exception as exc:
            print(f"  ERROR: {exc}", file=sys.stderr)


if __name__ == '__main__':
    main()
