Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
dns.py191 linesDownload Raw Back to layers
1import struct2import time3from dataclasses import dataclass4from typing import List5from typing import Literal6 7from mitmproxy import dns8from mitmproxy import flow as mflow9from mitmproxy.net.dns import response_codes10from mitmproxy.proxy import commands11from mitmproxy.proxy import events12from mitmproxy.proxy import layer13from mitmproxy.proxy.context import Context14from mitmproxy.proxy.utils import expect15 16_LENGTH_LABEL = struct.Struct("!H")17 18 19@dataclass20class DnsRequestHook(commands.StartHook):21    """22    A DNS query has been received.23    """24 25    flow: dns.DNSFlow26 27 28@dataclass29class DnsResponseHook(commands.StartHook):30    """31    A DNS response has been received or set.32    """33 34    flow: dns.DNSFlow35 36 37@dataclass38class DnsErrorHook(commands.StartHook):39    """40    A DNS error has occurred.41    """42 43    flow: dns.DNSFlow44 45 46def pack_message(47    message: dns.DNSMessage, transport_protocol: Literal["tcp", "udp"]48) -> bytes:49    packed = message.packed50    if transport_protocol == "tcp":51        return struct.pack("!H", len(packed)) + packed52    else:53        return packed54 55 56class DNSLayer(layer.Layer):57    """58    Layer that handles resolving DNS queries.59    """60 61    flows: dict[int, dns.DNSFlow]62    req_buf: bytearray63    resp_buf: bytearray64 65    def __init__(self, context: Context):66        super().__init__(context)67        self.flows = {}68        self.req_buf = bytearray()69        self.resp_buf = bytearray()70 71    def handle_request(72        self, flow: dns.DNSFlow, msg: dns.DNSMessage73    ) -> layer.CommandGenerator[None]:74        flow.request = msg  # if already set, continue and query upstream again75        yield DnsRequestHook(flow)76        if flow.response:77            yield from self.handle_response(flow, flow.response)78        elif flow.error:79            yield from self.handle_error(flow, flow.error.msg)80        elif not self.context.server.address:81            yield from self.handle_error(82                flow, "No hook has set a response and there is no upstream server."83            )84        else:85            if not self.context.server.connected:86                err = yield commands.OpenConnection(self.context.server)87                if err:88                    yield from self.handle_error(flow, str(err))89                    # cannot recover from this90                    return91            packed = pack_message(flow.request, flow.server_conn.transport_protocol)92            yield commands.SendData(self.context.server, packed)93 94    def handle_response(95        self, flow: dns.DNSFlow, msg: dns.DNSMessage96    ) -> layer.CommandGenerator[None]:97        flow.response = msg98        yield DnsResponseHook(flow)99        if flow.response:100            packed = pack_message(flow.response, flow.client_conn.transport_protocol)101            yield commands.SendData(self.context.client, packed)102 103    def handle_error(self, flow: dns.DNSFlow, err: str) -> layer.CommandGenerator[None]:104        flow.error = mflow.Error(err)105        yield DnsErrorHook(flow)106        servfail = flow.request.fail(response_codes.SERVFAIL)107        yield commands.SendData(108            self.context.client,109            pack_message(servfail, flow.client_conn.transport_protocol),110        )111 112    def unpack_message(self, data: bytes, from_client: bool) -> List[dns.DNSMessage]:113        msgs: List[dns.DNSMessage] = []114 115        buf = self.req_buf if from_client else self.resp_buf116 117        if self.context.client.transport_protocol == "udp":118            msgs.append(dns.DNSMessage.unpack(data, timestamp=time.time()))119        elif self.context.client.transport_protocol == "tcp":120            buf.extend(data)121            size = len(buf)122            offset = 0123 124            while True:125                if size - offset < _LENGTH_LABEL.size:126                    break127                (expected_size,) = _LENGTH_LABEL.unpack_from(buf, offset)128                offset += _LENGTH_LABEL.size129                if expected_size == 0:130                    raise struct.error("Message length field cannot be zero")131 132                if size - offset < expected_size:133                    offset -= _LENGTH_LABEL.size134                    break135 136                data = bytes(buf[offset : expected_size + offset])137                offset += expected_size138                msgs.append(dns.DNSMessage.unpack(data, timestamp=time.time()))139 140            del buf[:offset]141        return msgs142 143    @expect(events.Start)144    def state_start(self, _) -> layer.CommandGenerator[None]:145        self._handle_event = self.state_query146        yield from ()147 148    @expect(events.DataReceived, events.ConnectionClosed)149    def state_query(self, event: events.Event) -> layer.CommandGenerator[None]:150        assert isinstance(event, events.ConnectionEvent)151        from_client = event.connection is self.context.client152 153        if isinstance(event, events.DataReceived):154            msgs: List[dns.DNSMessage] = []155            try:156                msgs = self.unpack_message(event.data, from_client)157            except struct.error as e:158                yield commands.Log(f"{event.connection} sent an invalid message: {e}")159                yield commands.CloseConnection(event.connection)160                self._handle_event = self.state_done161            else:162                for msg in msgs:163                    try:164                        flow = self.flows[msg.id]165                    except KeyError:166                        flow = dns.DNSFlow(167                            self.context.client, self.context.server, live=True168                        )169                        self.flows[msg.id] = flow170                    if from_client:171                        yield from self.handle_request(flow, msg)172                    else:173                        yield from self.handle_response(flow, msg)174 175        elif isinstance(event, events.ConnectionClosed):176            other_conn = self.context.server if from_client else self.context.client177            if other_conn.connected:178                yield commands.CloseConnection(other_conn)179            self._handle_event = self.state_done180            for flow in self.flows.values():181                flow.live = False182 183        else:184            raise AssertionError(f"Unexpected event: {event}")185 186    @expect(events.DataReceived, events.ConnectionClosed)187    def state_done(self, _) -> layer.CommandGenerator[None]:188        yield from ()189 190    _handle_event = state_start191 
codekingpro/portable-devtools · Team Ai