#!/usr/bin/env python3
"""Download Witness PAC, then generate modified customer PAC."""

import re
from pathlib import Path
from urllib.error import HTTPError, URLError
from urllib.request import urlopen

PAC_URL = (
    "https://api.dev.witness.ai/v1/peas/pac/"
    "b9d5e9b5c73e2767628cc109e0ab09734e39e3f1ef56e2a65d105bc24488d2e0.pac"
    "?enableStunnel=true"
)
WITNESS_FILE = Path("witness.pac")
CUSTOMER_FILE = Path("proxy.pac")
MODIFIED_FILE = Path("modified.pac")
TIMEOUT_SECONDS = 30
WITNESS_PROXY_RETURN = "PROXY 127.0.0.1:9411; DIRECT"


def download_witness_pac() -> bool:
    try:
        with urlopen(PAC_URL, timeout=TIMEOUT_SECONDS) as response:
            pac_content = response.read()
    except HTTPError as exc:
        print(f"HTTP error while downloading PAC: {exc.code} {exc.reason}")
        return False
    except URLError as exc:
        print(f"Network error while downloading PAC: {exc.reason}")
        return False

    WITNESS_FILE.write_bytes(pac_content)
    print(f"Saved witness PAC to {WITNESS_FILE.resolve()}")
    return True


def extract_witness_domains(witness_text: str) -> list[str]:
    # Domain keys in witness PAC are represented as: 'domain.tld':1
    domain_pattern = r"""['"]([a-z0-9.-]+\.[a-z]{2,})['"]\s*:\s*1"""
    domains = {match.group(1).lower() for match in re.finditer(domain_pattern, witness_text, re.IGNORECASE)}
    return sorted(domains)


def build_witness_block(domains: list[str], indent: str) -> str:
    domain_entries = ",".join(f"'{domain}':1" for domain in domains)
    return (
        f"{indent}// WitnessAI\n"
        f"{indent}var witnessDomains = {{ {domain_entries} }};\n"
        f"{indent}var witnessHostNeedle = host.toLowerCase();\n"
        f"{indent}while (witnessHostNeedle) {{\n"
        f"{indent}    if (witnessDomains.hasOwnProperty(witnessHostNeedle)) {{\n"
        f"{indent}        return '{WITNESS_PROXY_RETURN}';\n"
        f"{indent}    }}\n"
        f"{indent}    var witnessNextDot = witnessHostNeedle.indexOf('.');\n"
        f"{indent}    if (witnessNextDot === -1) {{\n"
        f"{indent}        break;\n"
        f"{indent}    }}\n"
        f"{indent}    witnessHostNeedle = witnessHostNeedle.substring(witnessNextDot + 1);\n"
        f"{indent}}}\n\n"
    )


def inject_witness_block(customer_pac_text: str, block: str) -> str:
    direct_return_pattern = re.compile(r"return\s+['\"]DIRECT['\"]\s*;")
    direct_returns = list(direct_return_pattern.finditer(customer_pac_text))
    if not direct_returns:
        raise ValueError("Could not find a default DIRECT return statement in customer PAC.")

    target_match = direct_returns[-1]
    line_start = customer_pac_text.rfind("\n", 0, target_match.start())
    insertion_at = 0 if line_start == -1 else line_start + 1
    return customer_pac_text[:insertion_at] + block + customer_pac_text[insertion_at:]


def main() -> int:
    if not download_witness_pac():
        return 1

    if not CUSTOMER_FILE.exists():
        print(f"Customer PAC not found: {CUSTOMER_FILE.resolve()}")
        return 1

    witness_text = WITNESS_FILE.read_text(encoding="utf-8", errors="replace")
    customer_text = CUSTOMER_FILE.read_text(encoding="utf-8", errors="replace")
    domains = extract_witness_domains(witness_text)

    if not domains:
        print("No witness domains found in witness PAC. Not writing modified PAC.")
        return 1

    default_direct_pattern = re.compile(r"return\s+['\"]DIRECT['\"]\s*;")
    direct_returns = list(default_direct_pattern.finditer(customer_text))
    if not direct_returns:
        print("Could not find default DIRECT return in customer PAC.")
        return 1

    line_start = customer_text.rfind("\n", 0, direct_returns[-1].start())
    if line_start == -1:
        indent = ""
    else:
        line_text = customer_text[line_start + 1 : direct_returns[-1].start()]
        indent_match = re.match(r"[ \t]*", line_text)
        indent = indent_match.group(0) if indent_match else ""
    witness_block = build_witness_block(domains, indent)
    modified_text = inject_witness_block(customer_text, witness_block)
    MODIFIED_FILE.write_text(modified_text, encoding="utf-8")

    print(f"Extracted {len(domains)} witness domains.")
    print(f"Saved modified customer PAC to {MODIFIED_FILE.resolve()}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
