Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
mode_servers.py547 linesDownload Raw Back to proxy
1"""2This module defines "server instances", which manage3the TCP/UDP servers spawned by mitmproxy as specified by the proxy mode.4 5Example:6 7    mode = ProxyMode.parse("reverse:https://example.com")8    inst = ServerInstance.make(mode, manager_that_handles_callbacks)9    await inst.start()10    # TCP server is running now.11"""12 13from __future__ import annotations14 15import asyncio16import errno17import json18import logging19import os20import socket21import sys22import textwrap23import typing24from abc import ABCMeta25from abc import abstractmethod26from contextlib import contextmanager27from pathlib import Path28from typing import cast29from typing import ClassVar30from typing import Generic31from typing import get_args32from typing import TYPE_CHECKING33from typing import TypeVar34 35import mitmproxy_rs36from mitmproxy import ctx37from mitmproxy import flow38from mitmproxy import platform39from mitmproxy.connection import Address40from mitmproxy.net import local_ip41from mitmproxy.net.free_port import get_free_port42from mitmproxy.proxy import commands43from mitmproxy.proxy import layers44from mitmproxy.proxy import mode_specs45from mitmproxy.proxy import server46from mitmproxy.proxy.context import Context47from mitmproxy.proxy.layer import Layer48from mitmproxy.utils import human49 50if sys.version_info < (3, 11):51    from typing_extensions import Self  # pragma: no cover52else:53    from typing import Self54 55if TYPE_CHECKING:56    from mitmproxy.master import Master57 58logger = logging.getLogger(__name__)59 60 61class ProxyConnectionHandler(server.LiveConnectionHandler):62    master: Master63 64    def __init__(self, master, r, w, options, mode):65        self.master = master66        super().__init__(r, w, options, mode)67        self.log_prefix = f"{human.format_address(self.client.peername)}: "68 69    async def handle_hook(self, hook: commands.StartHook) -> None:70        with self.timeout_watchdog.disarm():71            # We currently only support single-argument hooks.72            (data,) = hook.args()73            await self.master.addons.handle_lifecycle(hook)74            if isinstance(data, flow.Flow):75                await data.wait_for_resume()  # pragma: no cover76 77 78M = TypeVar("M", bound=mode_specs.ProxyMode)79 80 81class ServerManager(typing.Protocol):82    # temporary workaround: for UDP, we use the 4-tuple because we don't have a uuid.83    connections: dict[tuple | str, ProxyConnectionHandler]84 85    @contextmanager86    def register_connection(87        self, connection_id: tuple | str, handler: ProxyConnectionHandler88    ): ...  # pragma: no cover89 90 91class ServerInstance(Generic[M], metaclass=ABCMeta):92    __modes: ClassVar[dict[str, type[ServerInstance]]] = {}93 94    last_exception: Exception | None = None95 96    def __init__(self, mode: M, manager: ServerManager):97        self.mode: M = mode98        self.manager: ServerManager = manager99 100    def __init_subclass__(cls, **kwargs):101        """Register all subclasses so that make() finds them."""102        # extract mode from Generic[Mode].103        mode = get_args(cls.__orig_bases__[0])[0]  # type: ignore104        if not isinstance(mode, TypeVar):105            assert issubclass(mode, mode_specs.ProxyMode)106            assert mode.type_name not in ServerInstance.__modes107            ServerInstance.__modes[mode.type_name] = cls108 109    @classmethod110    def make(111        cls,112        mode: mode_specs.ProxyMode | str,113        manager: ServerManager,114    ) -> Self:115        if isinstance(mode, str):116            mode = mode_specs.ProxyMode.parse(mode)117        inst = ServerInstance.__modes[mode.type_name](mode, manager)118 119        if not isinstance(inst, cls):120            raise ValueError(f"{mode!r} is not a spec for a {cls.__name__} server.")121 122        return inst123 124    @property125    @abstractmethod126    def is_running(self) -> bool:127        pass128 129    async def start(self) -> None:130        try:131            await self._start()132        except Exception as e:133            self.last_exception = e134            raise135        else:136            self.last_exception = None137        if self.listen_addrs:138            addrs = " and ".join({human.format_address(a) for a in self.listen_addrs})139            logger.info(f"{self.mode.description} listening at {addrs}.")140        else:141            logger.info(f"{self.mode.description} started.")142 143    async def stop(self) -> None:144        listen_addrs = self.listen_addrs145        try:146            await self._stop()147        except Exception as e:148            self.last_exception = e149            raise150        else:151            self.last_exception = None152        if listen_addrs:153            addrs = " and ".join({human.format_address(a) for a in listen_addrs})154            logger.info(f"{self.mode.description} at {addrs} stopped.")155        else:156            logger.info(f"{self.mode.description} stopped.")157 158    @abstractmethod159    async def _start(self) -> None:160        pass161 162    @abstractmethod163    async def _stop(self) -> None:164        pass165 166    @property167    @abstractmethod168    def listen_addrs(self) -> tuple[Address, ...]:169        pass170 171    @abstractmethod172    def make_top_layer(self, context: Context) -> Layer:173        pass174 175    def to_json(self) -> dict:176        return {177            "type": self.mode.type_name,178            "description": self.mode.description,179            "full_spec": self.mode.full_spec,180            "is_running": self.is_running,181            "last_exception": str(self.last_exception) if self.last_exception else None,182            "listen_addrs": self.listen_addrs,183        }184 185    async def handle_stream(186        self,187        reader: asyncio.StreamReader | mitmproxy_rs.Stream,188        writer: asyncio.StreamWriter | mitmproxy_rs.Stream | None = None,189    ) -> None:190        if writer is None:191            assert isinstance(reader, mitmproxy_rs.Stream)192            writer = reader193        handler = ProxyConnectionHandler(194            ctx.master, reader, writer, ctx.options, self.mode195        )196        handler.layer = self.make_top_layer(handler.layer.context)197        if isinstance(self.mode, mode_specs.TransparentMode):198            assert isinstance(writer, asyncio.StreamWriter)199            s = cast(socket.socket, writer.get_extra_info("socket"))200            try:201                assert platform.original_addr202                original_dst = platform.original_addr(s)203            except Exception as e:204                logger.error(f"Transparent mode failure: {e!r}")205                writer.close()206                return207            else:208                handler.layer.context.client.sockname = original_dst209                handler.layer.context.server.address = original_dst210        elif isinstance(211            self.mode,212            (mode_specs.WireGuardMode, mode_specs.LocalMode, mode_specs.TunMode),213        ):  # pragma: no cover on platforms without wg-test-client214            handler.layer.context.server.address = writer.get_extra_info(215                "remote_endpoint", handler.layer.context.client.sockname216            )217 218        with self.manager.register_connection(handler.layer.context.client.id, handler):219            await handler.handle_client()220 221 222class AsyncioServerInstance(ServerInstance[M], metaclass=ABCMeta):223    _servers: list[224        asyncio.Server225        | mitmproxy_rs.udp.UdpServer226        | mitmproxy_rs.wireguard.WireGuardServer227    ]228 229    def __init__(self, *args, **kwargs) -> None:230        self._servers = []231        super().__init__(*args, **kwargs)232 233    @property234    def is_running(self) -> bool:235        return bool(self._servers)236 237    @property238    def listen_addrs(self) -> tuple[Address, ...]:239        addrs = []240        for s in self._servers:241            if isinstance(242                s, (mitmproxy_rs.udp.UdpServer, mitmproxy_rs.wireguard.WireGuardServer)243            ):244                addrs.append(s.getsockname())245            else:246                try:247                    addrs.extend(sock.getsockname() for sock in s.sockets)248                except OSError:  # pragma: no cover249                    pass  # this can fail during shutdown, see https://github.com/mitmproxy/mitmproxy/issues/6529250        return tuple(addrs)251 252    async def _start(self) -> None:253        assert not self._servers254        host = self.mode.listen_host(ctx.options.listen_host)255        port = self.mode.listen_port(ctx.options.listen_port)256        assert port is not None257        try:258            self._servers = await self.listen(host, port)259        except OSError as e:260            message = f"{self.mode.description} failed to listen on {host or '*'}:{port} with {e}"261            if e.errno == errno.EADDRINUSE and self.mode.custom_listen_port is None:262                assert (263                    self.mode.custom_listen_host is None264                )  # since [@ [listen_addr:]listen_port]265                message += f"\nTry specifying a different port by using `--mode {self.mode.full_spec}@{port + 2}`."266            raise OSError(e.errno, message, e.filename) from e267 268    async def _stop(self) -> None:269        assert self._servers270        try:271            for s in self._servers:272                s.close()273            # https://github.com/python/cpython/issues/104344274            # await asyncio.gather(*[s.wait_closed() for s in self._servers])275        finally:276            # we always reset _server and ignore failures277            self._servers = []278 279    async def listen(280        self, host: str, port: int281    ) -> list[282        asyncio.Server283        | mitmproxy_rs.udp.UdpServer284        | mitmproxy_rs.wireguard.WireGuardServer285    ]:286        if self.mode.transport_protocol not in ("tcp", "udp", "both"):287            raise AssertionError(self.mode.transport_protocol)288 289        # workaround for https://github.com/python/cpython/issues/89856:290        # We want both IPv4 and IPv6 sockets to bind to the same port.291        # This may fail (https://github.com/mitmproxy/mitmproxy/pull/5542#issuecomment-1222803291),292        # so we try to cover the 99% case and then give up and fall back to what asyncio does.293        if port == 0:294            try:295                return await self.listen(host, get_free_port())296            except Exception as e:297                logger.debug(298                    f"Failed to listen on a single port ({e!r}), falling back to default behavior."299                )300 301        servers: list[302            asyncio.Server303            | mitmproxy_rs.udp.UdpServer304            | mitmproxy_rs.wireguard.WireGuardServer305        ] = []306        if self.mode.transport_protocol in ("tcp", "both"):307            servers.append(await asyncio.start_server(self.handle_stream, host, port))308        if self.mode.transport_protocol in ("udp", "both"):309            # we start two servers for dual-stack support.310            # On Linux, this would also be achievable by toggling IPV6_V6ONLY off, but this here works cross-platform.311            if host == "":312                ipv4 = await self.start_udp_based_server("0.0.0.0", port)313                servers.append(ipv4)314                try:315                    ipv6 = await self.start_udp_based_server(316                        "::", ipv4.getsockname()[1]317                    )318                    servers.append(ipv6)  # pragma: no cover319                except Exception:  # pragma: no cover320                    logger.debug("Failed to listen on '::', listening on IPv4 only.")321            else:322                servers.append(await self.start_udp_based_server(host, port))323 324        return servers325 326    async def start_udp_based_server(327        self, host, port328    ) -> mitmproxy_rs.udp.UdpServer | mitmproxy_rs.wireguard.WireGuardServer:329        return await mitmproxy_rs.udp.start_udp_server(330            host,331            port,332            self.handle_stream,333        )334 335 336class WireGuardServerInstance(AsyncioServerInstance[mode_specs.WireGuardMode]):337    server_key: str338    client_key: str339    pubkey: str340 341    def make_top_layer(342        self, context: Context343    ) -> Layer:  # pragma: no cover on platforms without wg-test-client344        return layers.modes.TransparentProxy(context)345 346    async def _start(self) -> None:347        if self.mode.data:348            conf_path = Path(self.mode.data).expanduser()349        else:350            conf_path = Path(ctx.options.confdir).expanduser() / "wireguard.conf"351 352        if not conf_path.exists():353            conf_path.parent.mkdir(parents=True, exist_ok=True)354            conf_path.write_text(355                json.dumps(356                    {357                        "server_key": mitmproxy_rs.wireguard.genkey(),358                        "client_key": mitmproxy_rs.wireguard.genkey(),359                    },360                    indent=4,361                )362            )363 364        try:365            c = json.loads(conf_path.read_text())366            self.server_key = c["server_key"]367            self.client_key = c["client_key"]368        except Exception as e:369            raise ValueError(f"Invalid configuration file ({conf_path}): {e}") from e370 371        # error early on invalid keys372        self.pubkey = mitmproxy_rs.wireguard.pubkey(self.client_key)373        _ = mitmproxy_rs.wireguard.pubkey(self.server_key)374 375        await super()._start()376 377        conf = self.client_conf()378        assert conf379        logger.info("-" * 60 + "\n" + conf + "\n" + "-" * 60)380 381    async def start_udp_based_server(382        self, host, port383    ) -> mitmproxy_rs.wireguard.WireGuardServer:384        return await mitmproxy_rs.wireguard.start_wireguard_server(385            host,386            port,387            self.server_key,388            [self.pubkey],389            self.handle_stream,390            self.handle_stream,391        )392 393    def client_conf(self) -> str | None:394        if not self._servers:395            return None396        host = (397            self.mode.listen_host(ctx.options.listen_host)398            or local_ip.get_local_ip()399            or local_ip.get_local_ip6()400        )401        port = self.mode.listen_port(ctx.options.listen_port)402        return textwrap.dedent(403            f"""404            [Interface]405            PrivateKey = {self.client_key}406            Address = 10.0.0.1/32407            DNS = 10.0.0.53408 409            [Peer]410            PublicKey = {mitmproxy_rs.wireguard.pubkey(self.server_key)}411            AllowedIPs = 0.0.0.0/0412            Endpoint = {host}:{port}413            """414        ).strip()415 416    def to_json(self) -> dict:417        return {"wireguard_conf": self.client_conf(), **super().to_json()}418 419 420class LocalRedirectorInstance(ServerInstance[mode_specs.LocalMode]):421    _server: ClassVar[mitmproxy_rs.local.LocalRedirector | None] = None422    """The local redirector daemon. Will be started once and then reused for all future instances."""423    _instance: ClassVar[LocalRedirectorInstance | None] = None424    """The current LocalRedirectorInstance. Will be unset again if an instance is stopped."""425    listen_addrs = ()426 427    @property428    def is_running(self) -> bool:429        return self._instance is not None430 431    def make_top_layer(self, context: Context) -> Layer:432        return layers.modes.TransparentProxy(context)433 434    @classmethod435    async def redirector_handle_stream(436        cls,437        stream: mitmproxy_rs.Stream,438    ) -> None:439        if cls._instance is not None:440            await cls._instance.handle_stream(stream)441 442    async def _start(self) -> None:443        if self._instance:444            raise RuntimeError("Cannot spawn more than one local redirector.")445 446        if self.mode.data:447            spec = f"{self.mode.data},!{os.getpid()}"448        else:449            spec = f"!{os.getpid()}"450 451        cls = self.__class__452        cls._instance = self  # assign before awaiting to avoid races453        if cls._server is None:454            try:455                cls._server = await mitmproxy_rs.local.start_local_redirector(456                    cls.redirector_handle_stream,457                    cls.redirector_handle_stream,458                )459            except Exception:460                cls._instance = None461                raise462 463        cls._server.set_intercept(spec)464 465    async def _stop(self) -> None:466        assert self._instance467        assert self._server468        self.__class__._instance = None469        # We're not shutting down the server because we want to avoid additional UAC prompts.470        self._server.set_intercept("")471 472 473class RegularInstance(AsyncioServerInstance[mode_specs.RegularMode]):474    def make_top_layer(self, context: Context) -> Layer:475        return layers.modes.HttpProxy(context)476 477 478class UpstreamInstance(AsyncioServerInstance[mode_specs.UpstreamMode]):479    def make_top_layer(self, context: Context) -> Layer:480        return layers.modes.HttpUpstreamProxy(context)481 482 483class TransparentInstance(AsyncioServerInstance[mode_specs.TransparentMode]):484    def make_top_layer(self, context: Context) -> Layer:485        return layers.modes.TransparentProxy(context)486 487 488class ReverseInstance(AsyncioServerInstance[mode_specs.ReverseMode]):489    def make_top_layer(self, context: Context) -> Layer:490        return layers.modes.ReverseProxy(context)491 492 493class Socks5Instance(AsyncioServerInstance[mode_specs.Socks5Mode]):494    def make_top_layer(self, context: Context) -> Layer:495        return layers.modes.Socks5Proxy(context)496 497 498class DnsInstance(AsyncioServerInstance[mode_specs.DnsMode]):499    def make_top_layer(self, context: Context) -> Layer:500        return layers.DNSLayer(context)501 502 503class TunInstance(ServerInstance[mode_specs.TunMode]):504    _server: mitmproxy_rs.tun.TunInterface | None = None505    listen_addrs = ()506 507    def make_top_layer(508        self, context: Context509    ) -> Layer:  # pragma: no cover mocked in tests510        return layers.modes.TransparentProxy(context)511 512    @property513    def is_running(self) -> bool:514        return self._server is not None515 516    @property517    def tun_name(self) -> str | None:518        if self._server:519            return self._server.tun_name()520        else:521            return None522 523    def to_json(self) -> dict:524        return {"tun_name": self.tun_name, **super().to_json()}525 526    async def _start(self) -> None:527        assert self._server is None528        self._server = await mitmproxy_rs.tun.create_tun_interface(529            self.handle_stream,530            self.handle_stream,531            tun_name=self.mode.data or None,532        )533        logger.info(f"TUN interface created: {self._server.tun_name()}")534 535    async def _stop(self) -> None:536        assert self._server is not None537        try:538            self._server.close()539            await self._server.wait_closed()540        finally:541            self._server = None542 543 544# class Http3Instance(AsyncioServerInstance[mode_specs.Http3Mode]):545#     def make_top_layer(self, context: Context) -> Layer:546#         return layers.modes.HttpProxy(context)547 
codekingpro/portable-devtools · Team Ai