codekingpro/portable-devtools
115k
1"""2Proxy Server Implementation using asyncio.3The very high level overview is as follows:4 5 - Spawn one coroutine per client connection and create a reverse proxy layer to example.com6 - Process any commands from layer (such as opening a server connection)7 - Wait for any IO and send it as events to top layer.8"""9 10import abc11import asyncio12import collections13import logging14import time15from collections.abc import Awaitable16from collections.abc import Callable17from collections.abc import MutableMapping18from contextlib import contextmanager19from dataclasses import dataclass20from types import TracebackType21from typing import Literal22 23from OpenSSL import SSL24 25import mitmproxy_rs26from mitmproxy import http27from mitmproxy import options as moptions28from mitmproxy import tls29from mitmproxy.connection import Address30from mitmproxy.connection import Client31from mitmproxy.connection import Connection32from mitmproxy.connection import ConnectionState33from mitmproxy.proxy import commands34from mitmproxy.proxy import events35from mitmproxy.proxy import layer36from mitmproxy.proxy import layers37from mitmproxy.proxy import mode_specs38from mitmproxy.proxy import server_hooks39from mitmproxy.proxy.context import Context40from mitmproxy.proxy.layers.http import HTTPMode41from mitmproxy.utils import asyncio_utils42from mitmproxy.utils import human43from mitmproxy.utils.data import pkg_data44 45logger = logging.getLogger(__name__)46 47UDP_TIMEOUT = 2048 49 50class TimeoutWatchdog:51 last_activity: float52 timeout: int53 can_timeout: asyncio.Event54 blocker: int55 56 def __init__(self, timeout: int, callback: Callable[[], Awaitable]):57 self.timeout = timeout58 self.callback = callback59 self.last_activity = time.time()60 self.can_timeout = asyncio.Event()61 self.can_timeout.set()62 self.blocker = 063 64 def register_activity(self):65 self.last_activity = time.time()66 67 async def watch(self):68 try:69 while True:70 await self.can_timeout.wait()71 await asyncio.sleep(self.timeout - (time.time() - self.last_activity))72 if self.last_activity + self.timeout < time.time():73 await self.callback()74 return75 except asyncio.CancelledError:76 return77 78 @contextmanager79 def disarm(self):80 self.can_timeout.clear()81 self.blocker += 182 try:83 yield84 finally:85 self.blocker -= 186 if self.blocker == 0:87 self.register_activity()88 self.can_timeout.set()89 90 91@dataclass92class ConnectionIO:93 handler: asyncio.Task | None = None94 reader: asyncio.StreamReader | mitmproxy_rs.Stream | None = None95 writer: asyncio.StreamWriter | mitmproxy_rs.Stream | None = None96 97 98class ConnectionHandler(metaclass=abc.ABCMeta):99 transports: MutableMapping[Connection, ConnectionIO]100 timeout_watchdog: TimeoutWatchdog101 client: Client102 max_conns: collections.defaultdict[Address, asyncio.Semaphore]103 layer: "layer.Layer"104 wakeup_timer: set[asyncio.Task]105 106 def __init__(self, context: Context) -> None:107 self.client = context.client108 self.transports = {}109 self.max_conns = collections.defaultdict(lambda: asyncio.Semaphore(5))110 self.wakeup_timer = set()111 112 # Ask for the first layer right away.113 # In a reverse proxy scenario, this is necessary as we would otherwise hang114 # on protocols that start with a server greeting.115 self.layer = layer.NextLayer(context, ask_on_start=True)116 if self.client.transport_protocol == "tcp":117 timeout = context.options.tcp_timeout118 else:119 timeout = UDP_TIMEOUT120 self.timeout_watchdog = TimeoutWatchdog(timeout, self.on_timeout)121 122 self._server_event_lock = asyncio.Lock()123 124 # workaround for https://bugs.python.org/issue40124 / https://bugs.python.org/issue29930125 self._drain_lock = asyncio.Lock()126 127 async def handle_client(self) -> None:128 asyncio_utils.set_current_task_debug_info(129 name=f"client handler",130 client=self.client.peername,131 )132 watch = asyncio_utils.create_task(133 self.timeout_watchdog.watch(),134 name="timeout watchdog",135 keep_ref=False,136 client=self.client.peername,137 )138 139 self.log("client connect")140 await self.handle_hook(server_hooks.ClientConnectedHook(self.client))141 if self.client.error:142 self.log("client kill connection")143 writer = self.transports.pop(self.client).writer144 assert writer145 writer.close()146 else:147 await self.server_event(events.Start())148 handler = asyncio_utils.create_task(149 self.handle_connection(self.client),150 name=f"client connection handler",151 keep_ref=False,152 client=self.client.peername,153 )154 self.transports[self.client].handler = handler155 await asyncio.wait([handler])156 if not handler.cancelled() and (e := handler.exception()):157 self.log(158 f"connection handler has crashed: {e}",159 logging.ERROR,160 exc_info=(type(e), e, e.__traceback__),161 )162 163 watch.cancel()164 while self.wakeup_timer:165 timer = self.wakeup_timer.pop()166 timer.cancel()167 168 self.log("client disconnect")169 self.client.timestamp_end = time.time()170 await self.handle_hook(server_hooks.ClientDisconnectedHook(self.client))171 172 if self.transports:173 self.log("closing transports...", logging.DEBUG)174 for io in self.transports.values():175 if io.handler:176 io.handler.cancel("client disconnected")177 await asyncio.wait(178 [x.handler for x in self.transports.values() if x.handler]179 )180 self.log("transports closed!", logging.DEBUG)181 182 async def open_connection(self, command: commands.OpenConnection) -> None:183 if not command.connection.address:184 self.log(f"Cannot open connection, no hostname given.")185 await self.server_event(186 events.OpenConnectionCompleted(187 command, f"Cannot open connection, no hostname given."188 )189 )190 return191 192 hook_data = server_hooks.ServerConnectionHookData(193 client=self.client, server=command.connection194 )195 await self.handle_hook(server_hooks.ServerConnectHook(hook_data))196 if err := command.connection.error:197 self.log(198 f"server connection to {human.format_address(command.connection.address)} killed before connect: {err}"199 )200 await self.handle_hook(server_hooks.ServerConnectErrorHook(hook_data))201 await self.server_event(202 events.OpenConnectionCompleted(command, f"Connection killed: {err}")203 )204 return205 206 async with self.max_conns[command.connection.address]:207 reader: asyncio.StreamReader | mitmproxy_rs.Stream208 writer: asyncio.StreamWriter | mitmproxy_rs.Stream209 try:210 command.connection.timestamp_start = time.time()211 if command.connection.transport_protocol == "tcp":212 reader, writer = await asyncio.open_connection(213 *command.connection.address,214 local_addr=command.connection.sockname,215 )216 elif command.connection.transport_protocol == "udp":217 reader = writer = await mitmproxy_rs.udp.open_udp_connection(218 *command.connection.address,219 local_addr=command.connection.sockname,220 )221 else:222 raise AssertionError(command.connection.transport_protocol)223 except (OSError, asyncio.CancelledError) as e:224 err = str(e)225 if not err: # str(CancelledError()) returns empty string.226 err = "connection cancelled"227 self.log(f"error establishing server connection: {err}")228 command.connection.error = err229 await self.handle_hook(server_hooks.ServerConnectErrorHook(hook_data))230 await self.server_event(events.OpenConnectionCompleted(command, err))231 if isinstance(e, asyncio.CancelledError):232 # From https://docs.python.org/3/library/asyncio-exceptions.html#asyncio.CancelledError:233 # > In almost all situations the exception must be re-raised.234 # It is not really defined what almost means here, but we play safe.235 raise236 else:237 if command.connection.transport_protocol == "tcp":238 # TODO: Rename to `timestamp_setup` and make it agnostic for both TCP (SYN/ACK) and UDP (DNS resl.)239 command.connection.timestamp_tcp_setup = time.time()240 command.connection.state = ConnectionState.OPEN241 command.connection.peername = writer.get_extra_info("peername")242 command.connection.sockname = writer.get_extra_info("sockname")243 self.transports[command.connection] = ConnectionIO(244 handler=asyncio.current_task(),245 reader=reader,246 writer=writer,247 )248 249 assert command.connection.peername250 if command.connection.address[0] != command.connection.peername[0]:251 addr = f"{human.format_address(command.connection.address)} ({human.format_address(command.connection.peername)})"252 else:253 addr = human.format_address(command.connection.address)254 self.log(f"server connect {addr}")255 await self.handle_hook(server_hooks.ServerConnectedHook(hook_data))256 await self.server_event(events.OpenConnectionCompleted(command, None))257 258 try:259 await self.handle_connection(command.connection)260 finally:261 self.log(f"server disconnect {addr}")262 command.connection.timestamp_end = time.time()263 await self.handle_hook(264 server_hooks.ServerDisconnectedHook(hook_data)265 )266 267 async def wakeup(self, request: commands.RequestWakeup) -> None:268 await asyncio.sleep(request.delay)269 task = asyncio.current_task()270 assert task is not None271 self.wakeup_timer.discard(task)272 await self.server_event(events.Wakeup(request))273 274 async def handle_connection(self, connection: Connection) -> None:275 """276 Handle a connection for its entire lifetime.277 This means we read until EOF,278 but then possibly also keep on waiting for our side of the connection to be closed.279 """280 cancelled = None281 reader = self.transports[connection].reader282 assert reader283 while True:284 try:285 data = await reader.read(65535)286 if not data:287 raise OSError("Connection closed by peer.")288 except OSError:289 break290 except asyncio.CancelledError as e:291 cancelled = e292 break293 294 await self.server_event(events.DataReceived(connection, data))295 296 try:297 await self.drain_writers()298 except asyncio.CancelledError as e:299 cancelled = e300 break301 302 if cancelled is None and connection.transport_protocol == "tcp":303 # TCP connections can be half-closed.304 connection.state &= ~ConnectionState.CAN_READ305 else:306 connection.state = ConnectionState.CLOSED307 308 await self.server_event(events.ConnectionClosed(connection))309 310 if connection.state is ConnectionState.CAN_WRITE:311 # we may still use this connection to *send* stuff,312 # even though the remote has closed their side of the connection.313 # to make this work we keep this task running and wait for cancellation.314 try:315 await asyncio.Event().wait()316 except asyncio.CancelledError as e:317 cancelled = e318 319 try:320 writer = self.transports[connection].writer321 assert writer322 writer.close()323 except OSError:324 pass325 self.transports.pop(connection)326 327 if cancelled:328 raise cancelled329 330 async def drain_writers(self):331 """332 Drain all writers to create some backpressure. We won't continue reading until there's space available in our333 write buffers, so if we cannot write fast enough our own read buffers run full and the TCP recv stream is throttled.334 """335 async with self._drain_lock:336 for transport in list(self.transports.values()):337 if transport.writer is not None:338 try:339 await transport.writer.drain()340 except OSError as e:341 if transport.handler is not None:342 transport.handler.cancel(f"Error sending data: {e}")343 344 async def on_timeout(self) -> None:345 try:346 handler = self.transports[self.client].handler347 except KeyError: # pragma: no cover348 # there is a super short window between connection close and watchdog cancellation349 pass350 else:351 if self.client.transport_protocol == "tcp":352 self.log(f"Closing connection due to inactivity: {self.client}")353 assert handler354 handler.cancel("timeout")355 356 async def hook_task(self, hook: commands.StartHook) -> None:357 await self.handle_hook(hook)358 if hook.blocking:359 await self.server_event(events.HookCompleted(hook))360 361 @abc.abstractmethod362 async def handle_hook(self, hook: commands.StartHook) -> None:363 pass364 365 def log(366 self,367 message: str,368 level: int = logging.INFO,369 exc_info: Literal[True]370 | tuple[type[BaseException], BaseException, TracebackType | None]371 | None = None,372 ) -> None:373 logger.log(374 level, message, extra={"client": self.client.peername}, exc_info=exc_info375 )376 377 async def server_event(self, event: events.Event) -> None:378 # server_event is supposed to be completely sync without any `await` that could pause execution.379 # However, create_task with an [eager task factory] will schedule tasks immediately,380 # which causes [reentrancy issues]. So we put the entire thing behind a lock.381 #382 # [eager task factory]: https://docs.python.org/3/library/asyncio-task.html#eager-task-factory383 # [reentrancy issues]: https://github.com/mitmproxy/mitmproxy/issues/7027.384 async with self._server_event_lock:385 # No `await` beyond this point.386 387 self.timeout_watchdog.register_activity()388 try:389 layer_commands = self.layer.handle_event(event)390 for command in layer_commands:391 if isinstance(command, commands.OpenConnection):392 assert command.connection not in self.transports393 handler = asyncio_utils.create_task(394 self.open_connection(command),395 name=f"server connection handler {command.connection.address}",396 keep_ref=False,397 client=self.client.peername,398 )399 self.transports[command.connection] = ConnectionIO(400 handler=handler401 )402 elif isinstance(command, commands.RequestWakeup):403 task = asyncio_utils.create_task(404 self.wakeup(command),405 name=f"wakeup timer ({command.delay:.1f}s)",406 keep_ref=False,407 client=self.client.peername,408 )409 assert task is not None410 self.wakeup_timer.add(task)411 elif (412 isinstance(command, commands.ConnectionCommand)413 and command.connection not in self.transports414 ):415 pass # The connection has already been closed.416 elif isinstance(command, commands.SendData):417 writer = self.transports[command.connection].writer418 assert writer419 if not writer.is_closing():420 writer.write(command.data)421 elif isinstance(command, commands.CloseTcpConnection):422 self.close_connection(command.connection, command.half_close)423 elif isinstance(command, commands.CloseConnection):424 self.close_connection(command.connection, False)425 elif isinstance(command, commands.StartHook):426 asyncio_utils.create_task(427 self.hook_task(command),428 name=f"handle_hook({command.name})",429 keep_ref=True,430 client=self.client.peername,431 )432 elif isinstance(command, commands.Log):433 self.log(command.message, command.level)434 else:435 raise RuntimeError(f"Unexpected command: {command}")436 except Exception:437 self.log(f"mitmproxy has crashed!", logging.ERROR, exc_info=True)438 439 def close_connection(440 self, connection: Connection, half_close: bool = False441 ) -> None:442 if half_close:443 if not connection.state & ConnectionState.CAN_WRITE:444 return445 self.log(f"half-closing {connection}", logging.DEBUG)446 try:447 writer = self.transports[connection].writer448 assert writer449 if not writer.is_closing():450 writer.write_eof()451 except OSError:452 # if we can't write to the socket anymore we presume it completely dead.453 connection.state = ConnectionState.CLOSED454 else:455 connection.state &= ~ConnectionState.CAN_WRITE456 else:457 connection.state = ConnectionState.CLOSED458 459 if connection.state is ConnectionState.CLOSED:460 handler = self.transports[connection].handler461 assert handler462 handler.cancel("closed by command")463 464 465class LiveConnectionHandler(ConnectionHandler, metaclass=abc.ABCMeta):466 def __init__(467 self,468 reader: asyncio.StreamReader | mitmproxy_rs.Stream,469 writer: asyncio.StreamWriter | mitmproxy_rs.Stream,470 options: moptions.Options,471 mode: mode_specs.ProxyMode,472 ) -> None:473 client = Client(474 transport_protocol=writer.get_extra_info("transport_protocol", "tcp"),475 peername=writer.get_extra_info("peername"),476 sockname=writer.get_extra_info("sockname"),477 timestamp_start=time.time(),478 proxy_mode=mode,479 state=ConnectionState.OPEN,480 )481 context = Context(client, options)482 super().__init__(context)483 self.transports[client] = ConnectionIO(484 handler=None, reader=reader, writer=writer485 )486 487 488class SimpleConnectionHandler(LiveConnectionHandler): # pragma: no cover489 """Simple handler that does not really process any hooks."""490 491 hook_handlers: dict[str, Callable]492 493 def __init__(self, reader, writer, options, mode, hook_handlers):494 super().__init__(reader, writer, options, mode)495 self.hook_handlers = hook_handlers496 497 async def handle_hook(self, hook: commands.StartHook) -> None:498 if hook.name in self.hook_handlers:499 self.hook_handlers[hook.name](*hook.args())500 501 502if __name__ == "__main__": # pragma: no cover503 # simple standalone implementation for testing.504 loop = asyncio.get_event_loop()505 506 opts = moptions.Options()507 # options duplicated here to simplify testing setup508 opts.add_option(509 "store_streamed_bodies",510 bool,511 False,512 "",513 )514 opts.add_option(515 "connection_strategy",516 str,517 "lazy",518 "Determine when server connections should be established.",519 choices=("eager", "lazy"),520 )521 opts.add_option(522 "keep_host_header",523 bool,524 False,525 """526 Reverse Proxy: Keep the original host header instead of rewriting it527 to the reverse proxy target.528 """,529 )530 531 async def handle(reader, writer):532 layer_stack = [533 # lambda ctx: layers.ServerTLSLayer(ctx),534 # lambda ctx: layers.HttpLayer(ctx, HTTPMode.regular),535 # lambda ctx: setattr(ctx.server, "tls", True) or layers.ServerTLSLayer(ctx),536 # lambda ctx: layers.ClientTLSLayer(ctx),537 lambda ctx: layers.modes.ReverseProxy(ctx),538 lambda ctx: layers.HttpLayer(ctx, HTTPMode.transparent),539 ]540 541 def next_layer(nl: layer.NextLayer):542 layr = layer_stack.pop(0)(nl.context)543 layr.debug = " " * len(nl.context.layers)544 nl.layer = layr545 546 def request(flow: http.HTTPFlow):547 if "cached" in flow.request.path:548 flow.response = http.Response.make(418, f"(cached) {flow.request.text}")549 if "toggle-tls" in flow.request.path:550 if flow.request.url.startswith("https://"):551 flow.request.url = flow.request.url.replace("https://", "http://")552 else:553 flow.request.url = flow.request.url.replace("http://", "https://")554 if "redirect" in flow.request.path:555 flow.request.host = "httpbin.org"556 557 def tls_start_client(tls_start: tls.TlsData):558 # INSECURE559 ssl_context = SSL.Context(SSL.SSLv23_METHOD)560 ssl_context.use_privatekey_file(561 pkg_data.path(562 "../test/mitmproxy/data/verificationcerts/trusted-leaf.key"563 )564 )565 ssl_context.use_certificate_chain_file(566 pkg_data.path(567 "../test/mitmproxy/data/verificationcerts/trusted-leaf.crt"568 )569 )570 tls_start.ssl_conn = SSL.Connection(ssl_context)571 tls_start.ssl_conn.set_accept_state()572 573 def tls_start_server(tls_start: tls.TlsData):574 # INSECURE575 ssl_context = SSL.Context(SSL.SSLv23_METHOD)576 tls_start.ssl_conn = SSL.Connection(ssl_context)577 tls_start.ssl_conn.set_connect_state()578 if tls_start.context.client.sni is not None:579 tls_start.ssl_conn.set_tlsext_host_name(580 tls_start.context.client.sni.encode()581 )582 583 await SimpleConnectionHandler(584 reader,585 writer,586 opts,587 mode_specs.ProxyMode.parse("reverse:http://127.0.0.1:3000/"),588 {589 "next_layer": next_layer,590 "request": request,591 "tls_start_client": tls_start_client,592 "tls_start_server": tls_start_server,593 },594 ).handle_client()595 596 coro = asyncio.start_server(handle, "127.0.0.1", 8080, loop=loop)597 server = loop.run_until_complete(coro)598 599 # Serve requests until Ctrl+C is pressed600 assert server.sockets601 print(f"Serving on {human.format_address(server.sockets[0].getsockname())}")602 try:603 loop.run_forever()604 except KeyboardInterrupt:605 pass606 607 # Close the server608 server.close()609 loop.run_until_complete(server.wait_closed())610 loop.close()611 