diff --git a/main.py b/main.py index d80add6..15cf6dd 100644 --- a/main.py +++ b/main.py @@ -12,6 +12,7 @@ import asyncio import json import logging import os +import subprocess import sys from cert_installer import install_ca, is_ca_trusted @@ -75,6 +76,70 @@ def parse_args(): return parser.parse_args() +def _windows_listener_details(host: str, port: int) -> tuple[str, str] | None: + """Best-effort lookup of the process listening on host:port on Windows.""" + if os.name != "nt": + return None + + try: + result = subprocess.run( + ["netstat", "-ano", "-p", "tcp"], + capture_output=True, + text=True, + check=True, + ) + except Exception: + return None + + port_suffix = f":{port}" + host_variants = { + f"{host}:{port}", + f"0.0.0.0:{port}", + f"[::]:{port}", + f"[::1]:{port}", + f"::{port}", + } + + pid = None + for line in result.stdout.splitlines(): + parts = line.split() + if len(parts) < 5 or parts[0] != "TCP": + continue + local_addr, state, candidate_pid = parts[1], parts[3].upper(), parts[4] + if state != "LISTENING": + continue + if local_addr in host_variants or local_addr.endswith(port_suffix): + pid = candidate_pid + break + + if not pid: + return None + + try: + proc = subprocess.run( + ["tasklist", "/FI", f"PID eq {pid}", "/FO", "CSV", "/NH"], + capture_output=True, + text=True, + check=True, + ) + line = proc.stdout.strip().splitlines()[0] + name = line.split(",")[0].strip('"') if line else "unknown" + except Exception: + name = "unknown" + + return pid, name + + +def _is_addr_in_use_error(exc: OSError) -> bool: + text = str(exc).lower() + return ( + getattr(exc, "errno", None) in {48, 98, 10048} + or getattr(exc, "winerror", None) == 10048 + or "address already in use" in text + or "only one usage of each socket address" in text + ) + + def main(): args = parse_args() config_path = args.config @@ -196,6 +261,20 @@ def main(): try: asyncio.run(ProxyServer(config).start()) + except OSError as e: + if _is_addr_in_use_error(e): + host = config.get("listen_host", "127.0.0.1") + port = config.get("listen_port", 8080) + log.error("Cannot listen on %s:%d because that address is already in use.", host, port) + details = _windows_listener_details(host, port) + if details: + pid, name = details + log.error("Port %d is currently held by PID %s (%s).", port, pid, name) + log.error( + "Stop the other process or choose another port, for example: python main.py -p 9090" + ) + sys.exit(1) + raise except KeyboardInterrupt: log.info("Stopped") diff --git a/mitm.py b/mitm.py index e3c91b1..1c6ed47 100644 --- a/mitm.py +++ b/mitm.py @@ -11,15 +11,17 @@ Requires: pip install cryptography """ import datetime +import ipaddress import logging import os +import re import ssl import tempfile from cryptography import x509 from cryptography.hazmat.primitives import hashes, serialization from cryptography.hazmat.primitives.asymmetric import rsa -from cryptography.x509.oid import NameOID +from cryptography.x509.oid import ExtendedKeyUsageOID, NameOID log = logging.getLogger("MITM") @@ -38,15 +40,19 @@ class MITMCertManager: def _ensure_ca(self): if os.path.exists(CA_KEY_FILE) and os.path.exists(CA_CERT_FILE): - with open(CA_KEY_FILE, "rb") as f: - self._ca_key = serialization.load_pem_private_key( - f.read(), password=None - ) - with open(CA_CERT_FILE, "rb") as f: - self._ca_cert = x509.load_pem_x509_certificate(f.read()) - log.info("Loaded CA from %s", CA_DIR) - else: - self._create_ca() + try: + with open(CA_KEY_FILE, "rb") as f: + self._ca_key = serialization.load_pem_private_key( + f.read(), password=None + ) + with open(CA_CERT_FILE, "rb") as f: + self._ca_cert = x509.load_pem_x509_certificate(f.read()) + log.info("Loaded CA from %s", CA_DIR) + return + except Exception as exc: + log.warning("Existing CA is unreadable, generating a new one: %s", exc) + + self._create_ca() def _create_ca(self): os.makedirs(CA_DIR, exist_ok=True) @@ -59,13 +65,14 @@ class MITMCertManager: x509.NameAttribute(NameOID.ORGANIZATION_NAME, "MasterHttpRelayVPN"), ]) now = datetime.datetime.now(datetime.timezone.utc) + ca_public_key = self._ca_key.public_key() self._ca_cert = ( x509.CertificateBuilder() .subject_name(subject) .issuer_name(issuer) - .public_key(self._ca_key.public_key()) + .public_key(ca_public_key) .serial_number(x509.random_serial_number()) - .not_valid_before(now) + .not_valid_before(now - datetime.timedelta(days=1)) .not_valid_after(now + datetime.timedelta(days=3650)) .add_extension( x509.BasicConstraints(ca=True, path_length=0), critical=True @@ -84,6 +91,14 @@ class MITMCertManager: ), critical=True, ) + .add_extension( + x509.SubjectKeyIdentifier.from_public_key(ca_public_key), + critical=False, + ) + .add_extension( + x509.AuthorityKeyIdentifier.from_issuer_public_key(ca_public_key), + critical=False, + ) .sign(self._ca_key, hashes.SHA256()) ) @@ -105,8 +120,9 @@ class MITMCertManager: if domain not in self._ctx_cache: key_pem, cert_pem = self._generate_domain_cert(domain) - cert_file = os.path.join(self._cert_dir, f"{domain}.crt") - key_file = os.path.join(self._cert_dir, f"{domain}.key") + cache_name = self._safe_cache_name(domain) + cert_file = os.path.join(self._cert_dir, f"{cache_name}.crt") + key_file = os.path.join(self._cert_dir, f"{cache_name}.key") ca_pem = self._ca_cert.public_bytes(serialization.Encoding.PEM) with open(cert_file, "wb") as f: @@ -122,23 +138,60 @@ class MITMCertManager: return self._ctx_cache[domain] def _generate_domain_cert(self, domain: str): + normalized_name, san_entries = self._build_subject_alt_names(domain) key = rsa.generate_private_key( public_exponent=65537, key_size=2048 ) + public_key = key.public_key() subject = x509.Name([ - x509.NameAttribute(NameOID.COMMON_NAME, domain), + x509.NameAttribute( + NameOID.COMMON_NAME, + normalized_name if len(normalized_name) <= 64 else "MasterHttpRelayVPN", + ), ]) now = datetime.datetime.now(datetime.timezone.utc) cert = ( x509.CertificateBuilder() .subject_name(subject) .issuer_name(self._ca_cert.subject) - .public_key(key.public_key()) + .public_key(public_key) .serial_number(x509.random_serial_number()) - .not_valid_before(now) - .not_valid_after(now + datetime.timedelta(days=365)) + .not_valid_before(now - datetime.timedelta(days=1)) + .not_valid_after(now + datetime.timedelta(days=90)) .add_extension( - x509.SubjectAlternativeName([x509.DNSName(domain)]), + x509.BasicConstraints(ca=False, path_length=None), + critical=True, + ) + .add_extension( + x509.KeyUsage( + digital_signature=True, + key_encipherment=True, + key_cert_sign=False, + crl_sign=False, + content_commitment=False, + data_encipherment=False, + key_agreement=False, + encipher_only=False, + decipher_only=False, + ), + critical=True, + ) + .add_extension( + x509.ExtendedKeyUsage([ExtendedKeyUsageOID.SERVER_AUTH]), + critical=False, + ) + .add_extension( + x509.SubjectAlternativeName(san_entries), + critical=False, + ) + .add_extension( + x509.SubjectKeyIdentifier.from_public_key(public_key), + critical=False, + ) + .add_extension( + x509.AuthorityKeyIdentifier.from_issuer_public_key( + self._ca_key.public_key() + ), critical=False, ) .sign(self._ca_key, hashes.SHA256()) @@ -151,3 +204,17 @@ class MITMCertManager: ) cert_pem = cert.public_bytes(serialization.Encoding.PEM) return key_pem, cert_pem + + @staticmethod + def _build_subject_alt_names(domain: str): + name = domain.strip().rstrip(".").strip("[]") + try: + ip = ipaddress.ip_address(name) + return name, [x509.IPAddress(ip)] + except ValueError: + normalized = name.encode("idna").decode("ascii") + return normalized, [x509.DNSName(normalized)] + + @staticmethod + def _safe_cache_name(domain: str) -> str: + return re.sub(r"[^A-Za-z0-9._-]", "_", domain)