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