[dns] minor fixes

This commit is contained in:
Manuel Meitinger 2022-03-27 15:23:59 +02:00
parent 4c2365ebeb
commit af2251c2ef
6 changed files with 80 additions and 62 deletions

View file

@ -8,7 +8,7 @@ from typing import List, Type
import mitmproxy.addons.next_layer # noqa
from mitmproxy import hooks, log, addonmanager
from mitmproxy.proxy import server_hooks, layer
from mitmproxy.proxy.layers import http, modes, tcp, tls, websocket
from mitmproxy.proxy.layers import dns, http, modes, tcp, tls, websocket
known = set()
@ -106,6 +106,16 @@ with outfile.open("w") as f, contextlib.redirect_stdout(f):
]
)
category(
"DNS",
"",
[
dns.DnsRequestHook,
dns.DnsResponseHook,
dns.DnsErrorHook,
]
)
category(
"TCP",
"",

View file

@ -61,7 +61,7 @@ class DnsServer:
"""Start a DNS server. Disabled by default."""
)
loader.add_option(
"dns_listen_host", Optional[str], "",
"dns_listen_host", str, "",
"""Address to bind DNS server to."""
)
loader.add_option(
@ -87,9 +87,7 @@ class DnsServer:
async def refresh_server(self) -> None:
async with self._lock:
if self.server:
await self.shutdown_server()
self.server = None
await self.shutdown_server()
if ctx.options.dns_server:
self.server = await udp.start_server(
self.handle_connection,
@ -101,18 +99,21 @@ class DnsServer:
ctx.log.info(f"DNS server listening at {' and '.join(addrs)}")
async def shutdown_server(self) -> None:
ctx.log.info("Stopping server...")
self.server.close()
await self.server.wait_closed()
self.server = None
if self.server is not None:
ctx.log.info("Stopping server...")
self.server.close()
await self.server.wait_closed()
self.server = None
async def handle_connection(self, r: asyncio.StreamReader, w: asyncio.StreamWriter) -> None:
peername = w.get_extra_info('peername')
asyncio_utils.set_task_debug_info(
asyncio.current_task(),
name=f"DnsServer.handle_connection",
client=peername,
)
current_task = asyncio.current_task()
if current_task is not None:
asyncio_utils.set_task_debug_info(
current_task,
name=f"DnsServer.handle_connection",
client=peername,
)
handler = DnsConnectionHandler(self.master, r, w, self.options)
self._connections[peername] = handler
try:
@ -123,7 +124,7 @@ class DnsServer:
def server_connect(self, ctx: server_hooks.ServerConnectionHookData) -> None:
assert ctx.server.address
self_connect = (
ctx.server.protocol == ConnectionProtocol.UDP
ctx.server.protocol is ConnectionProtocol.UDP
and
ctx.server.address[1] == self.options.dns_listen_port
and

View file

@ -347,7 +347,7 @@ class Dumper:
def dns_response(self, f: dns.DNSFlow):
# TODO this needs to be refined
if self.match(f) and f.response.answers:
if self.match(f) and f.response and f.response.answers:
self.echo("{server} answers {client} with {name} {ttl}s {class_} {type} {value}".format(
client=human.format_address(f.client_conn.peername),
server=human.format_address(f.server_conn.address),

View file

@ -397,24 +397,24 @@ class ResourceRecord(BypassInitStateObject):
return ".".join(labels)
@classmethod
def A(cls, name: str, ip: IPv4Address, *, ttl = DEFAULT_TTL) -> ResourceRecord:
def A(cls, name: str, ip: IPv4Address, *, ttl: int = DEFAULT_TTL) -> ResourceRecord:
"""Create an IPv4 resource record."""
return ResourceRecord(name, Type.A, Class.IN, ttl, ip.packed)
@classmethod
def AAAA(cls, name: str, ip: IPv6Address, *, ttl = DEFAULT_TTL) -> ResourceRecord:
def AAAA(cls, name: str, ip: IPv6Address, *, ttl: int = DEFAULT_TTL) -> ResourceRecord:
"""Create an IPv6 resource record."""
return ResourceRecord(name, Type.AAAA, Class.IN, ttl, ip.packed)
@classmethod
def CNAME(cls, alias: str, canonical: str, *, ttl = DEFAULT_TTL) -> ResourceRecord:
def CNAME(cls, alias: str, canonical: str, *, ttl: int = DEFAULT_TTL) -> ResourceRecord:
"""Create a canonical internet name resource record."""
return ResourceRecord(alias, Type.CNAME, Class.IN, ttl, ResourceRecord.encode_domain_name(canonical))
return ResourceRecord(alias, Type.CNAME, Class.IN, ttl, ResourceRecord.pack_domain_name(canonical))
@classmethod
def PTR(cls, inaddr: str, ptr: str, *, ttl = DEFAULT_TTL) -> ResourceRecord:
def PTR(cls, inaddr: str, ptr: str, *, ttl: int = DEFAULT_TTL) -> ResourceRecord:
"""Create a canonical internet name resource record."""
return ResourceRecord(inaddr, Type.PTR, Class.IN, ttl, ResourceRecord.encode_domain_name(ptr))
return ResourceRecord(inaddr, Type.PTR, Class.IN, ttl, ResourceRecord.pack_domain_name(ptr))
# comments are taken from rfc1035
@ -628,7 +628,7 @@ class Message(BypassInitStateObject):
except struct.error as e:
raise struct.error(f"question #{i}: {str(e)}")
def unpack_rrs(section: List[ResourceRecord], section_name: str, count: int) -> int:
def unpack_rrs(section: List[ResourceRecord], section_name: str, count: int) -> None:
nonlocal buffer, offset
for i in range(0, count):
try:

View file

@ -62,53 +62,60 @@ class DNSLayer(layer.Layer):
@classmethod
def simple_resolve(cls, questions: List[dns.Question]) -> List[dns.ResourceRecord]:
answers = []
answers: List[dns.ResourceRecord] = []
def resolve_by_name(family: socket.AddressFamily, ip: Callable[[str], Union[ipaddress.IPv4Address, ipaddress.IPv6Address]]) -> None:
nonlocal answers, question
try:
addrinfos = socket.getaddrinfo(host=question.name, port=0, family=family)
except socket.gaierror as e:
if e.errno == socket.EAI_NODATA:
for question in questions:
def resolve_by_name(
family: socket.AddressFamily,
ip: Callable[[str], Union[ipaddress.IPv4Address, ipaddress.IPv6Address]]
) -> None:
nonlocal answers, question
try:
addrinfos = socket.getaddrinfo(host=question.name, port=0, family=family)
except socket.gaierror as e:
if e.errno == socket.EAI_NODATA:
raise DnsResolveError(dns.ResponseCode.NXDOMAIN)
else:
# NOTE might fail on Windows for IPv6 queries:
# https://stackoverflow.com/questions/66755681/getaddrinfo-c-on-windows-not-handling-ipv6-correctly-returning-error-code-1
raise DnsResolveError(dns.ResponseCode.SERVFAIL)
for addrinfo in addrinfos:
_, _, _, _, addr = addrinfo
answers.append(dns.ResourceRecord(
name=question.name,
type=question.type,
class_=question.class_,
ttl=dns.ResourceRecord.DEFAULT_TTL,
data=ip(addr[0]).packed,
))
def resolve_by_addr(
suffix: str,
ip: Callable[[List[str]], Union[ipaddress.IPv4Address, ipaddress.IPv6Address]]
) -> bool:
nonlocal answers, question
if not question.name.lower().endswith(suffix.lower()):
return False
try:
addr = ip(question.name[0:-len(suffix)].split(".")[::-1])
except ValueError:
raise DnsResolveError(dns.ResponseCode.FORMERR)
try:
name, _, _ = socket.gethostbyaddr(str(addr))
except socket.herror:
raise DnsResolveError(dns.ResponseCode.NXDOMAIN)
else:
# NOTE might fail on Windows for IPv6 queries:
# https://stackoverflow.com/questions/66755681/getaddrinfo-c-on-windows-not-handling-ipv6-correctly-returning-error-code-1
except socket.gaierror:
raise DnsResolveError(dns.ResponseCode.SERVFAIL)
for addrinfo in addrinfos:
_, _, _, _, (addr, _) = addrinfo
answers.append(dns.ResourceRecord(
name=question.name,
type=question.type,
class_=question.class_,
ttl=dns.ResourceRecord.DEFAULT_TTL,
data=ip(addr).packed,
data=dns.ResourceRecord.pack_domain_name(name),
))
return True
def resolve_by_addr(suffix: str, ip: Callable[[List[str]], Union[ipaddress.IPv4Address, ipaddress.IPv6Address]]) -> bool:
nonlocal answers, question
if not question.name.lower().endswith(suffix.lower()):
return False
try:
addr = ip(question.name[0:-len(suffix)].split(".")[::-1])
except ValueError:
raise DnsResolveError(dns.ResponseCode.FORMERR)
try:
name, _, _ = socket.gethostbyaddr(str(addr))
except socket.herror:
raise DnsResolveError(dns.ResponseCode.NXDOMAIN)
except socket.gaierror:
raise DnsResolveError(dns.ResponseCode.SERVFAIL)
answers.append(dns.ResourceRecord(
name=question.name,
type=question.type,
class_=question.class_,
ttl=dns.ResourceRecord.DEFAULT_TTL,
data=dns.ResourceRecord.pack_domain_name(name),
))
return True
for question in questions:
if question.class_ is not dns.Class.IN:
raise DnsResolveError(dns.ResponseCode.NOTIMP)
if question.type is dns.Type.A:

View file

@ -154,9 +154,9 @@ class ConnectionHandler(metaclass=abc.ABCMeta):
try:
command.connection.timestamp_start = time.time()
open_connection = (
asyncio.open_connection if command.connection.protocol == ConnectionProtocol.TCP
asyncio.open_connection if command.connection.protocol is ConnectionProtocol.TCP
else
udp.open_connection if command.connection.protocol == ConnectionProtocol.UDP
udp.open_connection if command.connection.protocol is ConnectionProtocol.UDP
else
None
)