codekingpro/portable-devtools
115k
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 