#!/usr/bin/env python3
import argparse
import random
import re
import socket
import struct
import sys

DNS_SERVER = "8.8.8.8"
DNS_PORT = 53
DEFAULT_TIMEOUT = 2.0

QWERTY_NEIGHBORS = {
    "q": "wa", "w": "qeas", "e": "wrds", "r": "etdf", "t": "ryfg",
    "y": "tugh", "u": "yihj", "i": "uojk", "o": "ipkl", "p": "ol",
    "a": "qwsz", "s": "awedxz", "d": "serfxc", "f": "drtgcv", "g": "ftyhvb",
    "h": "gyujbn", "j": "huikmn", "k": "jiolm", "l": "kop",
    "z": "asx", "x": "zsdc", "c": "xdfv", "v": "cfgb", "b": "vghn",
    "n": "bhjm", "m": "njk",
}

HOMOGLYPH_MAP = {
    "o": ["0", "о", "ο"],   # CYRILLIC O, GREEK OMICRON
    "a": ["а"],                   # CYRILLIC A
    "e": ["е"],                   # CYRILLIC IE
    "i": ["1", "l", "і"],         # CYRILLIC BYELORUSSIAN-UKRAINIAN I
    "l": ["1", "i"],
    "c": ["с"],                   # CYRILLIC ES
    "p": ["р"],                   # CYRILLIC ER
    "x": ["х"],                   # CYRILLIC HA
    "y": ["у"],                   # CYRILLIC U
}

SUBSTRING_HOMOGLYPHS = [
    ("rn", "m"), ("m", "rn"),
    ("vv", "w"), ("w", "vv"),
    ("cl", "d"),
]

DOMAIN_RE = re.compile(r"^[A-Za-z0-9-]{1,63}(\.[A-Za-z0-9-]{1,63})+$")


def validate_domain(domain):
    domain = domain.strip()
    if not DOMAIN_RE.match(domain):
        raise ValueError(f"'{domain}' is not a valid domain")
    label, _, tld = domain.rpartition(".")
    if not label or label.startswith("-") or label.endswith("-"):
        raise ValueError(f"'{domain}' has an invalid second-level label")
    return label, tld


def generate_omissions(label):
    return {label[:i] + label[i + 1:] for i in range(len(label))} - {""}


def generate_repetitions(label):
    return {label[:i + 1] + label[i] + label[i + 1:] for i in range(len(label))}


def generate_adjacent_substitutions(label):
    candidates = set()
    for i, ch in enumerate(label):
        for neighbor in QWERTY_NEIGHBORS.get(ch, ""):
            candidates.add(label[:i] + neighbor + label[i + 1:])
    return candidates


def generate_bounded_insertions(label):
    candidates = set()
    n = len(label)
    for i in range(n + 1):
        before = label[i - 1] if i > 0 else ""
        after = label[i] if i < n else ""
        letters = set(QWERTY_NEIGHBORS.get(before, "")) | set(QWERTY_NEIGHBORS.get(after, ""))
        for letter in letters:
            candidates.add(label[:i] + letter + label[i:])
    return candidates


def generate_homoglyphs(label):
    candidates = set()
    for i, ch in enumerate(label):
        for glyph in HOMOGLYPH_MAP.get(ch, []):
            candidates.add(label[:i] + glyph + label[i + 1:])
    for old, new in SUBSTRING_HOMOGLYPHS:
        start = 0
        while True:
            idx = label.find(old, start)
            if idx == -1:
                break
            candidates.add(label[:idx] + new + label[idx + len(old):])
            start = idx + 1
    return candidates


def generate_all_candidates(label, tld):
    generators = [
        generate_omissions,
        generate_repetitions,
        generate_adjacent_substitutions,
        generate_bounded_insertions,
        generate_homoglyphs,
    ]
    candidates = {}
    for gen in generators:
        for mutated in gen(label):
            if mutated and mutated != label:
                candidates[f"{mutated}.{tld}"] = gen.__name__
    return candidates


def domain_to_ascii(domain):
    try:
        return domain.encode("idna").decode("ascii")
    except UnicodeError:
        return None


def encode_qname(ascii_domain):
    out = bytearray()
    for part in ascii_domain.split("."):
        if not part or len(part) > 63:
            raise ValueError(f"invalid label '{part}' in '{ascii_domain}'")
        out.append(len(part))
        out.extend(part.encode("ascii"))
    out.append(0)
    if len(out) > 255:
        raise ValueError(f"encoded name too long: '{ascii_domain}'")
    return bytes(out)


def build_dns_query(ascii_domain, query_id):
    header = struct.pack(">HHHHHH", query_id, 0x0100, 1, 0, 0, 0)
    question = encode_qname(ascii_domain) + struct.pack(">HH", 1, 1)
    return header + question


def decode_name(msg, offset):
    labels = []
    jumps = 0
    while True:
        if offset >= len(msg):
            raise ValueError("truncated DNS message while reading name")
        length = msg[offset]
        if length & 0xC0 == 0xC0:
            jumps += 1
            if jumps > 20:
                raise ValueError("too many DNS compression pointer jumps")
            if offset + 1 >= len(msg):
                raise ValueError("truncated compression pointer")
            pointer = ((length & 0x3F) << 8) | msg[offset + 1]
            if jumps == 1:
                end_offset = offset + 2
            offset = pointer
            continue
        if length == 0:
            offset += 1
            if not labels:
                return "", (end_offset if jumps else offset)
            return ".".join(labels), (end_offset if jumps else offset)
        offset += 1
        if offset + length > len(msg):
            raise ValueError("truncated DNS message while reading label")
        labels.append(msg[offset:offset + length].decode("ascii", errors="replace"))
        offset += length


def parse_dns_response(msg, expected_id):
    if len(msg) < 12:
        raise ValueError("DNS response too short")
    id_, flags = struct.unpack(">HH", msg[0:4])
    qdcount, ancount, _, _ = struct.unpack(">HHHH", msg[4:12])
    if id_ != expected_id:
        return -1, []
    rcode = flags & 0x000F

    offset = 12
    for _ in range(qdcount):
        _, offset = decode_name(msg, offset)
        offset += 4  # QTYPE + QCLASS

    ips = []
    for _ in range(ancount):
        _, offset = decode_name(msg, offset)
        if offset + 10 > len(msg):
            raise ValueError("truncated DNS resource record")
        rtype, _, _, rdlength = struct.unpack(">HHIH", msg[offset:offset + 10])
        offset += 10
        if offset + rdlength > len(msg):
            raise ValueError("truncated DNS resource record data")
        rdata = msg[offset:offset + rdlength]
        if rtype == 1 and rdlength == 4:
            ips.append(socket.inet_ntoa(rdata))
        offset += rdlength

    return rcode, ips


def resolve_a_record(ascii_domain, server, port, timeout):
    query_id = random.randint(0, 65535)
    try:
        packet = build_dns_query(ascii_domain, query_id)
    except ValueError:
        return []

    try:
        with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as sock:
            sock.settimeout(timeout)
            sock.sendto(packet, (server, port))
            data, _ = sock.recvfrom(512)
    except (socket.timeout, OSError):
        return []

    try:
        rcode, ips = parse_dns_response(data, query_id)
    except (ValueError, struct.error):
        return []

    if rcode != 0:
        return []
    return ips


def parse_args():
    parser = argparse.ArgumentParser(
        description="Generate typosquat domain candidates and check which resolve via DNS."
    )
    parser.add_argument("domain", help="root domain to generate typosquats for, e.g. example.com")
    parser.add_argument("--server", default=DNS_SERVER, help="DNS server to query (default: 8.8.8.8)")
    parser.add_argument("--port", type=int, default=DNS_PORT, help="DNS server port (default: 53)")
    parser.add_argument("--timeout", type=float, default=DEFAULT_TIMEOUT, help="per-query timeout in seconds")
    return parser.parse_args()


def main():
    args = parse_args()

    try:
        label, tld = validate_domain(args.domain)
    except ValueError as exc:
        print(f"error: {exc}", file=sys.stderr)
        return 1

    candidates = generate_all_candidates(label, tld)

    for full_domain in candidates:
        ascii_domain = domain_to_ascii(full_domain)
        if ascii_domain is None:
            continue

        try:
            ips = resolve_a_record(ascii_domain, args.server, args.port, args.timeout)
        except Exception as exc:
            print(f"warning: failed to check '{full_domain}': {exc}", file=sys.stderr)
            continue

        if ips:
            if full_domain != ascii_domain:
                print(f"{full_domain} ({ascii_domain}) -> {', '.join(ips)}")
            else:
                print(f"{full_domain} -> {', '.join(ips)}")

    return 0


if __name__ == "__main__":
    sys.exit(main())
