diff --git a/mitmproxy/addons/proxyserver.py b/mitmproxy/addons/proxyserver.py index 14fa47951..e69da4381 100644 --- a/mitmproxy/addons/proxyserver.py +++ b/mitmproxy/addons/proxyserver.py @@ -73,7 +73,7 @@ class Proxyserver: ) yield "DNS", self.dns_server, lambda x: setattr(self, "dns_server", x), ctx.options.dns_server, lambda: udp.start_server( self.handle_dns_datagram, - self.options.dns_listen_host or '127.0.0.1', + self.options.dns_listen_host or "127.0.0.1", self.options.dns_listen_port, transparent=self.options.dns_mode == "transparent" ) @@ -237,7 +237,7 @@ class Proxyserver: del self._connections[connection_id] async def handle_tcp_connection(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: - connection_id = ("tcp", writer.get_extra_info('peername'), writer.get_extra_info('sockname')) + connection_id = ("tcp", writer.get_extra_info("peername"), writer.get_extra_info("sockname")) self._connections[connection_id] = ProxyConnectionHandler(self.master, reader, writer, self.options) await self.handle_connection(connection_id) @@ -301,7 +301,8 @@ class Proxyserver: async def server_connect(self, ctx: server_hooks.ServerConnectionHookData): assert ctx.server.address - addrinfos = await asyncio.get_running_loop().getaddrinfo(*ctx.server.address[:2], proto=ctx.server.protocol.value) + host, port = ctx.server.address[:2] + addrinfos = await asyncio.get_running_loop().getaddrinfo(host, port, proto=ctx.server.protocol.value) for srv in self.running_servers: for sock in srv.sockets: for family, _, proto, _, addr in addrinfos: diff --git a/mitmproxy/net/udp.py b/mitmproxy/net/udp.py index 7c72583a1..1e06649ba 100644 --- a/mitmproxy/net/udp.py +++ b/mitmproxy/net/udp.py @@ -89,7 +89,7 @@ class DrainableDatagramProtocol(asyncio.DatagramProtocol): self._sock = None def __repr__(self) -> str: - return f'<{self.__class__.__name__} socket={self._sock!r}>' + return f"<{self.__class__.__name__} socket={self._sock!r}>" @property def sockets(self) -> Tuple[socket.socket, ...]: diff --git a/test/mitmproxy/addons/test_proxyserver.py b/test/mitmproxy/addons/test_proxyserver.py index e9b431e3b..3bc726651 100644 --- a/test/mitmproxy/addons/test_proxyserver.py +++ b/test/mitmproxy/addons/test_proxyserver.py @@ -234,34 +234,37 @@ async def test_dns_simple() -> None: flow = tdnsflow(resp=False) ps = Proxyserver() with taddons.context(ps) as tctx: - tctx.configure(ps, server=False, dns_server=True, dns_listen_port=5353, dns_mode="simple") + tctx.configure(ps, server=False, dns_server=True, dns_listen_host="127.0.0.1", dns_listen_port=0, dns_mode="simple") await ps.running() await tctx.master.await_log("DNS server listening at", level="info") await ps.dns_request(flow) assert flow.response await ps.shutdown_server() + await tctx.master.await_log("Stopping DNS server", level="info") async def test_dns_not_simple() -> None: flow = tdnsflow(resp=False) ps = Proxyserver() with taddons.context(ps) as tctx: - tctx.configure(ps, server=False, dns_server=True, dns_listen_port=5354, dns_mode="custom") + tctx.configure(ps, server=False, dns_server=True, dns_listen_host="127.0.0.1", dns_listen_port=0, dns_mode="custom") await ps.running() await tctx.master.await_log("DNS server listening at", level="info") await ps.dns_request(flow) assert not flow.response await ps.shutdown_server() + await tctx.master.await_log("Stopping DNS server", level="info") async def test_dns() -> None: ps = Proxyserver() with taddons.context(ps) as tctx: - tctx.configure(ps, server=False, dns_server=True, dns_listen_port=5355, dns_mode="simple") + tctx.configure(ps, server=False, dns_server=True, dns_listen_host="127.0.0.1", dns_listen_port=0, dns_mode="simple") await ps.running() await tctx.master.await_log("DNS server listening at", level="info") assert ps.dns_server - r, w = await udp.open_connection(*ps.dns_server.sockets[0].getsockname()[:2]) + dns_addr = ps.dns_server.sockets[0].getsockname()[:2] + r, w = await udp.open_connection(*dns_addr) req = tdnsreq() w.write(req.packed) resp = dns.Message.unpack(await r.read(udp.MAX_DATAGRAM_SIZE)) @@ -273,3 +276,4 @@ async def test_dns() -> None: assert req.id == resp.id and "8.8.8.8" in str(resp) assert len(ps._connections) == 1 await ps.shutdown_server() + await tctx.master.await_log("Stopping DNS server", level="info")