import os import socket import struct import tempfile import threading import unittest import netprobe def answer(ident, name, rcode=0, addresses=(), cname=None): header = struct.pack(">HHHHHH", ident, 0x8180 | rcode, 1, len(addresses) + (1 if cname else 0), 0, 0) question = netprobe.build_query(name, ident)[12:] records = b"" if cname: target = b"".join(bytes([len(l)]) + l.encode() for l in cname.split(".")) + b"\x00" records += b"\xc0\x0c" + struct.pack(">HHIH", 5, 1, 60, len(target)) + target for address in addresses: records += b"\xc0\x0c" + struct.pack(">HHIH", 1, 1, 60, 4) + socket.inet_aton(address) return header + question + records class ParseResponseTest(unittest.TestCase): def test_addresses(self): data = answer(7, "broker.lan", addresses=["192.168.0.10", "192.168.0.11"]) parsed = netprobe.parse_response(data, 7) self.assertEqual(parsed["rcode"], "NOERROR") self.assertEqual(parsed["addresses"], ["192.168.0.10", "192.168.0.11"]) def test_nxdomain(self): parsed = netprobe.parse_response(answer(1, "nope.invalid", rcode=3), 1) self.assertEqual(parsed["rcode"], "NXDOMAIN") self.assertEqual(parsed["addresses"], []) def test_cname(self): parsed = netprobe.parse_response(answer(2, "ha.example", cname="real.example", addresses=["10.0.0.1"]), 2) self.assertEqual(parsed["cnames"], ["real.example"]) self.assertEqual(parsed["addresses"], ["10.0.0.1"]) def test_id_mismatch(self): with self.assertRaises(ValueError): netprobe.parse_response(answer(3, "x.y"), 4) class DnsQueryTest(unittest.TestCase): def test_against_fake_server(self): server = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) server.bind(("127.0.0.1", 0)) port = server.getsockname()[1] def serve(): data, addr = server.recvfrom(512) ident = struct.unpack(">H", data[:2])[0] server.sendto(answer(ident, "broker.lan", addresses=["192.168.0.10"]), addr) server.close() threading.Thread(target=serve, daemon=True).start() result = netprobe.dns_query("127.0.0.1", "broker.lan", timeout=2, port=port) self.assertEqual(result["addresses"], ["192.168.0.10"]) self.assertEqual(result["rcode"], "NOERROR") class ResolvConfTest(unittest.TestCase): def test_parse(self): with tempfile.NamedTemporaryFile("w", delete=False) as fh: fh.write("# comment\nnameserver 192.168.0.1\nnameserver 8.8.8.8 # public\nsearch lan home\noptions ndots:1\n") try: parsed = netprobe.read_resolv_conf(fh.name) finally: os.unlink(fh.name) self.assertEqual(parsed["nameservers"], ["192.168.0.1", "8.8.8.8"]) self.assertEqual(parsed["search"], ["lan", "home"]) self.assertEqual(parsed["options"], ["ndots:1"]) def test_missing(self): parsed = netprobe.read_resolv_conf("/nonexistent/resolv.conf") self.assertEqual(parsed["nameservers"], []) self.assertIn("error", parsed) def test_hosts(self): with tempfile.NamedTemporaryFile("w", delete=False) as fh: fh.write("127.0.0.1 localhost\n192.168.0.5 ha.local homeassistant\n") try: self.assertEqual(netprobe.hosts_entries("homeassistant", fh.name), ["192.168.0.5"]) self.assertEqual(netprobe.hosts_entries("other", fh.name), []) finally: os.unlink(fh.name) class TcpTest(unittest.TestCase): def test_connect_ok_and_refused(self): listener = socket.socket() listener.bind(("127.0.0.1", 0)) listener.listen(1) port = listener.getsockname()[1] try: ok = netprobe.tcp_connect("127.0.0.1", port, 2) finally: listener.close() self.assertTrue(ok["ok"]) refused = netprobe.tcp_connect("127.0.0.1", port, 2) self.assertFalse(refused["ok"]) self.assertIn("error", refused) class SummaryTest(unittest.TestCase): def test_disagreement_flagged(self): result = { "host": "ha.milans.cloud", "port": 1883, "hosts_file": [], "resolv_conf": {"nameservers": ["192.168.0.1", "8.8.8.8"]}, "getaddrinfo": {"addresses": [], "error": "Name does not resolve"}, "nameservers": [ {"server": "192.168.0.1", "rcode": "NOERROR", "addresses": ["192.168.0.20"]}, {"server": "8.8.8.8", "rcode": "NXDOMAIN", "addresses": []}, ], "tcp": [{"address": "192.168.0.20", "ok": True, "ms": 2.0}], } lines = netprobe.summarize(result) self.assertTrue(any("disagree" in line for line in lines)) self.assertTrue(any("System resolver failed" in line for line in lines)) self.assertTrue(any("connected in" in line for line in lines)) if __name__ == "__main__": unittest.main()