Include server information in bind/listen errors. (#5495)

* include host and port in listen errors

* common UDP server and reuse message

* update tests

* fix reuse test for master
This commit is contained in:
Manuel Meitinger 2022-08-02 19:00:34 +02:00 committed by GitHub
parent 73c38f698d
commit fbcffee87f
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
5 changed files with 196 additions and 140 deletions

View file

@ -2,6 +2,8 @@
## Unreleased: mitmproxy next
* Include server information in bind/listen errors.
([#5495](https://github.com/mitmproxy/mitmproxy/pull/5495), @meitinger)
* Include information about lazy connection_strategy in related errors.
([#5465](https://github.com/mitmproxy/mitmproxy/pull/5465), @meitinger, @mhils)
* Add UDP layer and flow support.

View file

@ -12,6 +12,7 @@ Example:
from __future__ import annotations
import asyncio
import errno
import struct
import typing
from abc import ABCMeta, abstractmethod
@ -103,8 +104,49 @@ class ServerInstance(Generic[M], metaclass=ABCMeta):
pass
class TcpServerInstance(ServerInstance[M], metaclass=ABCMeta):
server: asyncio.Server | None = None
class AsyncioServerInstance(ServerInstance[M], metaclass=ABCMeta):
server: asyncio.Server | udp.UdpServer | None = None
async def start(self):
assert not self.server
host = self.mode.listen_host(ctx.options.listen_host)
port = self.mode.listen_port(ctx.options.listen_port)
try:
self.server = await self.listen(host, port)
except OSError as e:
message = f"{self.log_desc} failed to listen on {host or '*'}:{port} with {e}"
if e.errno == errno.EADDRINUSE and self.mode.custom_listen_port is None:
assert self.mode.custom_listen_host is None # since [@ [listen_addr:]listen_port]
message += f"\nTry specifying a different port by using `--mode {self.mode.full_spec}@{port + 1}`."
raise OSError(e.errno, message, e.filename) from e
addrs = {f"{human.format_address(s)}" for s in self.listen_addrs}
ctx.log.info(
f"{self.log_desc} listening at {' and '.join(addrs)}."
)
async def stop(self):
assert self.server
self.server.close()
await self.server.wait_closed()
ctx.log.info(f"Stopped {self.mode.type} proxy server.")
@abstractmethod
async def listen(self, host: str, port: int) -> asyncio.Server | udp.UdpServer:
pass
@property
@abstractmethod
def log_desc(self) -> str:
pass
@cached_property
def listen_addrs(self) -> tuple[Address, ...]:
assert self.server
return tuple(s.getsockname() for s in self.server.sockets)
class TcpServerInstance(AsyncioServerInstance[M], metaclass=ABCMeta):
@abstractmethod
def make_top_layer(self, context: Context) -> Layer:
@ -127,35 +169,13 @@ class TcpServerInstance(ServerInstance[M], metaclass=ABCMeta):
with self.manager.register_connection(connection_id, handler):
await handler.handle_client()
async def start(self):
assert not self.server
self.server = await asyncio.start_server(
async def listen(self, host: str, port: int) -> asyncio.Server:
return await asyncio.start_server(
self.handle_tcp_connection,
self.mode.listen_host(ctx.options.listen_host),
self.mode.listen_port(ctx.options.listen_port),
host,
port,
)
addrs = {f"{human.format_address(s)}" for s in self.listen_addrs}
ctx.log.info(
f"{self.log_desc} listening at {' and '.join(addrs)}."
)
@property
@abstractmethod
def log_desc(self) -> str:
pass
async def stop(self):
assert self.server
self.server.close()
await self.server.wait_closed()
ctx.log.info(f"Stopped {self.mode.type} proxy server.")
@cached_property
def listen_addrs(self) -> tuple[Address, ...]:
assert self.server
return tuple(s.getsockname() for s in self.server.sockets)
class RegularInstance(TcpServerInstance[mode_specs.RegularMode]):
log_desc = "HTTP(S) proxy"
@ -194,129 +214,106 @@ class Socks5Instance(TcpServerInstance[mode_specs.Socks5Mode]):
return layers.modes.Socks5Proxy(context)
class DnsInstance(ServerInstance[mode_specs.DnsMode]):
server: udp.UdpServer | None = None
class UdpServerInstance(AsyncioServerInstance[M], metaclass=ABCMeta):
async def start(self):
assert not self.server
self.server = await udp.start_server(
self.handle_dns_datagram,
self.mode.listen_host(ctx.options.listen_host),
self.mode.listen_port(ctx.options.listen_port),
@abstractmethod
def make_top_layer(self, context: Context) -> Layer:
pass
@abstractmethod
def make_connection_id(
self,
transport: asyncio.DatagramTransport,
data: bytes,
remote_addr: Address,
local_addr: Address,
) -> tuple | None:
pass
async def listen(self, host: str, port: int) -> udp.UdpServer:
return await udp.start_server(
self.handle_udp_datagram,
host,
port,
transparent=False
)
addrs = {f"{human.format_address(s)}" for s in self.listen_addrs}
ctx.log.info(
f"DNS server listening at {' and '.join(addrs)}."
)
async def stop(self):
assert self.server
self.server.close()
await self.server.wait_closed()
ctx.log.info(f"Stopped {self.mode.type} proxy server.")
def handle_dns_datagram(
def handle_udp_datagram(
self,
transport: asyncio.DatagramTransport,
data: bytes,
remote_addr: Address,
local_addr: Address,
) -> None:
connection_id = self.make_connection_id(transport, data, remote_addr, local_addr)
if connection_id is None:
return
if connection_id not in self.manager.connections:
reader = udp.DatagramReader()
writer = udp.DatagramWriter(transport, remote_addr, reader)
handler = ProxyConnectionHandler(
ctx.master, reader, writer, ctx.options, self.mode
)
handler.timeout_watchdog.CONNECTION_TIMEOUT = 20
handler.layer = self.make_top_layer(handler.layer.context)
# pre-register here - we may get datagrams before the task is executed.
self.manager.connections[connection_id] = handler
asyncio.create_task(self.handle_udp_connection(connection_id, handler))
else:
handler = self.manager.connections[connection_id]
reader = cast(udp.DatagramReader, handler.transports[handler.client].reader)
reader.feed_data(data, remote_addr)
async def handle_udp_connection(self, connection_id: tuple, handler: ProxyConnectionHandler) -> None:
with self.manager.register_connection(connection_id, handler):
await handler.handle_client()
class DnsInstance(UdpServerInstance[mode_specs.DnsMode]):
log_desc = "DNS server"
def make_top_layer(self, context: Context) -> Layer:
layer = layers.DNSLayer(context)
layer.context.server.address = (self.mode.data or "resolve-local", 53)
layer.context.server.transport_protocol = "udp"
return layer
def make_connection_id(
self,
transport: asyncio.DatagramTransport,
data: bytes,
remote_addr: Address,
local_addr: Address,
) -> tuple | None:
try:
dns_id = struct.unpack_from("!H", data, 0)
except struct.error:
ctx.log.info(
f"Invalid DNS datagram received from {human.format_address(remote_addr)}."
)
return
connection_id = ("udp", dns_id, remote_addr, local_addr)
if connection_id not in self.manager.connections:
reader = udp.DatagramReader()
writer = udp.DatagramWriter(transport, remote_addr, reader)
handler = ProxyConnectionHandler(
ctx.master, reader, writer, ctx.options, self.mode
)
handler.timeout_watchdog.CONNECTION_TIMEOUT = 20
handler.layer = layers.DNSLayer(handler.layer.context)
handler.layer.context.server.address = (self.mode.data or "resolve-local", 53)
handler.layer.context.server.transport_protocol = "udp"
# pre-register here - we may get datagrams before the task is executed.
self.manager.connections[connection_id] = handler
asyncio.create_task(self.handle_dns_connection(connection_id, handler))
return None
else:
handler = self.manager.connections[connection_id]
reader = cast(udp.DatagramReader, handler.transports[handler.client].reader)
reader.feed_data(data, remote_addr)
async def handle_dns_connection(self, connection_id, handler):
with self.manager.register_connection(connection_id, handler):
await handler.handle_client()
@cached_property
def listen_addrs(self) -> tuple[Address, ...]:
assert self.server
return tuple(s.getsockname() for s in self.server.sockets)
return ("udp", dns_id, remote_addr, local_addr)
class DtlsInstance(ServerInstance[mode_specs.DtlsMode]):
server: udp.UdpServer | None = None
class DtlsInstance(UdpServerInstance[mode_specs.DtlsMode]):
log_desc = "DTLS server"
async def start(self):
assert not self.server
self.server = await udp.start_server(
self.handle_dtls_datagram,
self.mode.listen_host(ctx.options.listen_host),
self.mode.listen_port(ctx.options.listen_port),
transparent=False
)
addrs = {f"{human.format_address(s)}" for s in self.listen_addrs}
ctx.log.info(
f"DTLS server listening at {' and '.join(addrs)}."
)
def make_top_layer(self, context: Context) -> Layer:
context.client.transport_protocol = "udp"
layer = layers.ServerTLSLayer(context)
layer.child_layer = layers.ClientTLSLayer(layer.context)
layer.child_layer.child_layer = layers.UDPLayer(layer.context)
layer.context.server.address = self.mode.address
layer.context.server.transport_protocol = "udp"
return layer
async def stop(self):
assert self.server
self.server.close()
await self.server.wait_closed()
ctx.log.info(f"Stopped {self.mode.type} proxy server.")
def handle_dtls_datagram(
self,
transport: asyncio.DatagramTransport,
data: bytes,
remote_addr: Address,
local_addr: Address,
):
connection_id = ("dtls", remote_addr, local_addr)
if connection_id not in self.manager.connections:
reader = udp.DatagramReader()
writer = udp.DatagramWriter(transport, remote_addr, reader)
handler = ProxyConnectionHandler(
ctx.master, reader, writer, ctx.options, self.mode
)
handler.timeout_watchdog.CONNECTION_TIMEOUT = 20
handler.layer.context.client.transport_protocol = "udp"
handler.layer = layers.ServerTLSLayer(handler.layer.context)
handler.layer.child_layer = layers.ClientTLSLayer(handler.layer.context)
handler.layer.child_layer.child_layer = layers.UDPLayer(handler.layer.context)
handler.layer.context.server.address = self.mode.address
handler.layer.context.server.transport_protocol = "udp"
# pre-register here - we may get datagrams before the task is executed.
self.manager.connections[connection_id] = handler
asyncio.create_task(self.handle_dtls_connection(connection_id, handler))
else:
handler = self.manager.connections[connection_id]
reader = cast(udp.DatagramReader, handler.transports[handler.client].reader)
reader.feed_data(data, remote_addr)
async def handle_dtls_connection(self, connection_id, handler): # pragma: no cover
with self.manager.register_connection(connection_id, handler):
await handler.handle_client()
@cached_property
def listen_addrs(self) -> tuple[Address, ...]:
assert self.server
return tuple(s.getsockname() for s in self.server.sockets)
def make_connection_id(
self,
transport: asyncio.DatagramTransport,
data: bytes,
remote_addr: Address,
local_addr: Address,
) -> tuple | None:
return ("dtls", remote_addr, local_addr)

View file

@ -1,3 +1,4 @@
import errno
import tornado.httpserver
import tornado.ioloop
@ -80,7 +81,13 @@ class WebMaster(master.Master):
# Add our web app.
http_server = tornado.httpserver.HTTPServer(self.app)
http_server.listen(self.options.web_port, self.options.web_host)
try:
http_server.listen(self.options.web_port, self.options.web_host)
except OSError as e:
message = f"Web server failed to listen on {self.options.web_host or '*'}:{self.options.web_port} with {e}"
if e.errno == errno.EADDRINUSE:
message += f"\nTry specifying a different port by using `--set web_port={self.options.web_port + 1}`."
raise OSError(e.errno, message, e.filename) from e
self.log.info(
f"Web server listening at http://{self.options.web_host}:{self.options.web_port}/",

View file

@ -2,6 +2,8 @@ import asyncio
from typing import cast
from unittest.mock import AsyncMock, MagicMock, Mock
import pytest
from mitmproxy.net import udp
from mitmproxy.proxy.mode_servers import DnsInstance, ServerInstance, DtlsInstance
from mitmproxy.test import taddons
@ -39,6 +41,24 @@ async def test_tcp_start_stop():
assert await tctx.master.await_log("Stopped regular proxy server.")
async def test_tcp_start_error():
manager = MagicMock()
with taddons.context() as tctx:
inst = ServerInstance.make("regular@127.0.0.1:0", manager)
await inst.start()
assert await tctx.master.await_log("proxy listening")
port = inst.listen_addrs[0][1]
inst2 = ServerInstance.make(f"regular@127.0.0.1:{port}", manager)
with pytest.raises(OSError, match=f"proxy failed to listen on 127\\.0\\.0\\.1:{port}"):
await inst2.start()
tctx.options.listen_host = "127.0.0.1"
tctx.options.listen_port = port
inst3 = ServerInstance.make(f"regular", manager)
with pytest.raises(OSError):
await inst3.start()
async def test_udp_start_stop():
manager = MagicMock()
@ -61,6 +81,19 @@ async def test_udp_start_stop():
assert await tctx.master.await_log("Stopped")
async def test_udp_start_error():
manager = MagicMock()
with taddons.context() as tctx:
inst = ServerInstance.make("dns@127.0.0.1:0", manager)
await inst.start()
assert await tctx.master.await_log("server listening")
port = inst.listen_addrs[0][1]
inst2 = ServerInstance.make(f"dns@127.0.0.1:{port}", manager)
with pytest.raises(OSError, match=f"server failed to listen on 127\\.0\\.0\\.1:{port}"):
await inst2.start()
async def test_dtls_start_stop(monkeypatch):
manager = MagicMock()
@ -82,12 +115,12 @@ async def test_udp_connection_reuse(monkeypatch):
manager.connections = {}
monkeypatch.setattr(udp, "DatagramWriter", MagicMock())
monkeypatch.setattr(DnsInstance, "handle_dns_connection", AsyncMock())
monkeypatch.setattr(DnsInstance, "handle_udp_connection", AsyncMock())
with taddons.context():
inst = cast(DnsInstance, ServerInstance.make("dns", manager))
inst.handle_dns_datagram(MagicMock(), b"\x00\x00\x01", ("remoteaddr", 0), ("localaddr", 0))
inst.handle_dns_datagram(MagicMock(), b"\x00\x00\x02", ("remoteaddr", 0), ("localaddr", 0))
inst.handle_udp_datagram(MagicMock(), b"\x00\x00\x01", ("remoteaddr", 0), ("localaddr", 0))
inst.handle_udp_datagram(MagicMock(), b"\x00\x00\x02", ("remoteaddr", 0), ("localaddr", 0))
await asyncio.sleep(0)
assert len(inst.manager.connections) == 1
@ -98,12 +131,12 @@ async def test_dtls_connection_reuse(monkeypatch):
manager.connections = {}
monkeypatch.setattr(udp, "DatagramWriter", MagicMock())
monkeypatch.setattr(DtlsInstance, "handle_dtls_connection", AsyncMock())
monkeypatch.setattr(DtlsInstance, "handle_udp_connection", AsyncMock())
with taddons.context():
inst = cast(DtlsInstance, ServerInstance.make("dtls:reverse:127.0.0.1:0", manager))
inst.handle_dtls_datagram(MagicMock(), b"\x00\x00\x01", ("remoteaddr", 0), ("localaddr", 0))
inst.handle_dtls_datagram(MagicMock(), b"\x00\x00\x02", ("remoteaddr", 0), ("localaddr", 0))
inst.handle_udp_datagram(MagicMock(), b"\x00\x00\x01", ("remoteaddr", 0), ("localaddr", 0))
inst.handle_udp_datagram(MagicMock(), b"\x00\x00\x02", ("remoteaddr", 0), ("localaddr", 0))
await asyncio.sleep(0)
assert len(inst.manager.connections) == 1

View file

@ -0,0 +1,17 @@
import asyncio
from unittest.mock import MagicMock
import pytest
from mitmproxy.options import Options
from mitmproxy.tools.web.master import WebMaster
async def test_reuse():
server = await asyncio.start_server(MagicMock(), host="127.0.0.1", port=0)
port = server.sockets[0].getsockname()[1]
master = WebMaster(Options())
master.options.web_host = "127.0.0.1"
master.options.web_port = port
with pytest.raises(OSError, match=f"--set web_port={port + 1}"):
await master.running()
server.close()