Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
websocket.py273 linesDownload Raw Back to layers
1import time2from collections.abc import Iterator3from dataclasses import dataclass4 5import wsproto.extensions6import wsproto.frame_protocol7import wsproto.utilities8from wsproto import ConnectionState9from wsproto.frame_protocol import Opcode10 11from mitmproxy import connection12from mitmproxy import http13from mitmproxy import websocket14from mitmproxy.proxy import commands15from mitmproxy.proxy import events16from mitmproxy.proxy import layer17from mitmproxy.proxy.commands import StartHook18from mitmproxy.proxy.context import Context19from mitmproxy.proxy.events import MessageInjected20from mitmproxy.proxy.utils import expect21 22 23@dataclass24class WebsocketStartHook(StartHook):25    """26    A WebSocket connection has commenced.27    """28 29    flow: http.HTTPFlow30 31 32@dataclass33class WebsocketMessageHook(StartHook):34    """35    Called when a WebSocket message is received from the client or36    server. The most recent message will be flow.messages[-1]. The37    message is user-modifiable. Currently there are two types of38    messages, corresponding to the BINARY and TEXT frame types.39    """40 41    flow: http.HTTPFlow42 43 44@dataclass45class WebsocketEndHook(StartHook):46    """47    A WebSocket connection has ended.48    You can check `flow.websocket.close_code` to determine why it ended.49    """50 51    flow: http.HTTPFlow52 53 54class WebSocketMessageInjected(MessageInjected[websocket.WebSocketMessage]):55    """56    The user has injected a custom WebSocket message.57    """58 59 60class WebsocketConnection(wsproto.Connection):61    """62    A very thin wrapper around wsproto.Connection:63 64     - we keep the underlying connection as an attribute for easy access.65     - we add a framebuffer for incomplete messages66     - we wrap .send() so that we can directly yield it.67    """68 69    conn: connection.Connection70    frame_buf: list[bytes]71 72    def __init__(self, *args, conn: connection.Connection, **kwargs):73        super().__init__(*args, **kwargs)74        self.conn = conn75        self.frame_buf = [b""]76 77    def send2(self, event: wsproto.events.Event) -> commands.SendData:78        data = self.send(event)79        return commands.SendData(self.conn, data)80 81    def __repr__(self):82        return f"WebsocketConnection<{self.state.name}, {self.conn}>"83 84 85class WebsocketLayer(layer.Layer):86    """87    WebSocket layer that intercepts and relays messages.88    """89 90    flow: http.HTTPFlow91    client_ws: WebsocketConnection92    server_ws: WebsocketConnection93 94    def __init__(self, context: Context, flow: http.HTTPFlow):95        super().__init__(context)96        self.flow = flow97 98    @expect(events.Start)99    def start(self, _) -> layer.CommandGenerator[None]:100        client_extensions = []101        server_extensions = []102 103        # Parse extension headers. We only support deflate at the moment and ignore everything else.104        assert self.flow.response  # satisfy type checker105        ext_header = self.flow.response.headers.get("Sec-WebSocket-Extensions", "")106        if ext_header:107            for ext in wsproto.utilities.split_comma_header(108                ext_header.encode("ascii", "replace")109            ):110                ext_name = ext.split(";", 1)[0].strip()111                if ext_name == wsproto.extensions.PerMessageDeflate.name:112                    client_deflate = wsproto.extensions.PerMessageDeflate()113                    client_deflate.finalize(ext)114                    client_extensions.append(client_deflate)115                    server_deflate = wsproto.extensions.PerMessageDeflate()116                    server_deflate.finalize(ext)117                    server_extensions.append(server_deflate)118                else:119                    yield commands.Log(120                        f"Ignoring unknown WebSocket extension {ext_name!r}."121                    )122 123        self.client_ws = WebsocketConnection(124            wsproto.ConnectionType.SERVER, client_extensions, conn=self.context.client125        )126        self.server_ws = WebsocketConnection(127            wsproto.ConnectionType.CLIENT, server_extensions, conn=self.context.server128        )129 130        yield WebsocketStartHook(self.flow)131 132        self._handle_event = self.relay_messages133 134    _handle_event = start135 136    @expect(events.DataReceived, events.ConnectionClosed, WebSocketMessageInjected)137    def relay_messages(self, event: events.Event) -> layer.CommandGenerator[None]:138        assert self.flow.websocket  # satisfy type checker139 140        if isinstance(event, events.ConnectionEvent):141            from_client = event.connection == self.context.client142            injected = False143        elif isinstance(event, WebSocketMessageInjected):144            from_client = event.message.from_client145            injected = True146        else:147            raise AssertionError(f"Unexpected event: {event}")148 149        from_str = "client" if from_client else "server"150        if from_client:151            src_ws = self.client_ws152            dst_ws = self.server_ws153        else:154            src_ws = self.server_ws155            dst_ws = self.client_ws156 157        if isinstance(event, events.DataReceived):158            src_ws.receive_data(event.data)159        elif isinstance(event, events.ConnectionClosed):160            src_ws.receive_data(None)161        elif isinstance(event, WebSocketMessageInjected):162            fragmentizer = Fragmentizer([], event.message.type == Opcode.TEXT)163            src_ws._events.extend(fragmentizer(event.message.content))164        else:  # pragma: no cover165            raise AssertionError(f"Unexpected event: {event}")166 167        for ws_event in src_ws.events():168            if isinstance(ws_event, wsproto.events.Message):169                is_text = isinstance(ws_event.data, str)170                if is_text:171                    typ = Opcode.TEXT172                    src_ws.frame_buf[-1] += ws_event.data.encode()173                else:174                    typ = Opcode.BINARY175                    src_ws.frame_buf[-1] += ws_event.data176 177                if ws_event.message_finished:178                    content = b"".join(src_ws.frame_buf)179 180                    fragmentizer = Fragmentizer(src_ws.frame_buf, is_text)181                    src_ws.frame_buf = [b""]182 183                    message = websocket.WebSocketMessage(184                        typ, from_client, content, injected=injected185                    )186                    self.flow.websocket.messages.append(message)187                    yield WebsocketMessageHook(self.flow)188 189                    if not message.dropped:190                        for msg in fragmentizer(message.content):191                            yield dst_ws.send2(msg)192 193                elif ws_event.frame_finished:194                    src_ws.frame_buf.append(b"")195 196            elif isinstance(ws_event, (wsproto.events.Ping, wsproto.events.Pong)):197                yield commands.Log(198                    f"Received WebSocket {ws_event.__class__.__name__.lower()} from {from_str} "199                    f"(payload: {bytes(ws_event.payload)!r})"200                )201                yield dst_ws.send2(ws_event)202            elif isinstance(ws_event, wsproto.events.CloseConnection):203                self.flow.websocket.timestamp_end = time.time()204                self.flow.websocket.closed_by_client = from_client205                self.flow.websocket.close_code = ws_event.code206                self.flow.websocket.close_reason = ws_event.reason207 208                for ws in [self.server_ws, self.client_ws]:209                    if ws.state in {210                        ConnectionState.OPEN,211                        ConnectionState.REMOTE_CLOSING,212                    }:213                        # response == original event, so no need to differentiate here.214                        yield ws.send2(ws_event)215                    yield commands.CloseConnection(ws.conn)216                yield WebsocketEndHook(self.flow)217                self.flow.live = False218                self._handle_event = self.done219            else:  # pragma: no cover220                raise AssertionError(f"Unexpected WebSocket event: {ws_event}")221 222    @expect(events.DataReceived, events.ConnectionClosed, WebSocketMessageInjected)223    def done(self, _) -> layer.CommandGenerator[None]:224        yield from ()225 226 227class Fragmentizer:228    """229    Theory (RFC 6455):230       Unless specified otherwise by an extension, frames have no semantic231       meaning.  An intermediary might coalesce and/or split frames, [...]232 233    Practice:234        Some WebSocket servers reject large payload sizes.235        Other WebSocket servers reject CONTINUATION frames.236 237    As a workaround, we either retain the original chunking or, if the payload has been modified, use ~4kB chunks.238    If one deals with web servers that do not support CONTINUATION frames, addons need to monkeypatch FRAGMENT_SIZE239    if they need to modify the message.240    """241 242    # A bit less than 4kb to accommodate for headers.243    FRAGMENT_SIZE = 4000244 245    def __init__(self, fragments: list[bytes], is_text: bool):246        self.fragment_lengths = [len(x) for x in fragments]247        self.is_text = is_text248 249    def msg(self, data: bytes, message_finished: bool):250        if self.is_text:251            data_str = data.decode(errors="replace")252            return wsproto.events.TextMessage(253                data_str, message_finished=message_finished254            )255        else:256            return wsproto.events.BytesMessage(data, message_finished=message_finished)257 258    def __call__(self, content: bytes) -> Iterator[wsproto.events.Message]:259        if len(content) == sum(self.fragment_lengths):260            # message has the same length, we can reuse the same sizes261            offset = 0262            for fl in self.fragment_lengths[:-1]:263                yield self.msg(content[offset : offset + fl], False)264                offset += fl265            yield self.msg(content[offset:], True)266        else:267            offset = 0268            total = len(content) - self.FRAGMENT_SIZE269            while offset < total:270                yield self.msg(content[offset : offset + self.FRAGMENT_SIZE], False)271                offset += self.FRAGMENT_SIZE272            yield self.msg(content[offset:], True)273 
codekingpro/portable-devtools · Team Ai