Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
server.py611 linesDownload Raw Back to proxy
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 
codekingpro/portable-devtools · Team Ai