diff --git a/netinfra/mdns/hap-bridge.py b/netinfra/mdns/hap-bridge.py index 2563e23..ee995a9 100644 --- a/netinfra/mdns/hap-bridge.py +++ b/netinfra/mdns/hap-bridge.py @@ -19,13 +19,14 @@ Design notes: """ import json +import os import logging import socket import struct import threading import time -CONFIG_PATH = "/usr/local/etc/hap-bridge.conf" +CONFIG_PATH = os.environ.get("CONFIG_PATH", "/usr/local/etc/hap-bridge.conf") MDNS_GROUP = "224.0.0.251" MDNS_PORT = 5353 REFRESH_INTERVAL = 60 @@ -60,6 +61,24 @@ def encode_name(name): return out + b"\x00" +def decode_rr_header(data, off): + """Return (rtype, rclass, ttl, rdata_len) for the RR header at offset.""" + return struct.unpack(">HHIH", data[off:off + 10]) + + +def merge_records(prev_recs, new_recs): + """Union record sets by (name, type); new records replace same-key old ones. + + Devices often answer a multi-question probe with partial record sets + (ecobee returns only the PTR, #619); accumulate across probes so a thin + response never discards previously learned SRV/TXT/A records. + """ + merged = {(r["name"], r["type"]): r for r in prev_recs} + for r in new_recs: + merged[(r["name"], r["type"])] = r + return list(merged.values()) + + def probe_device(ip, timeout=2.0): """Return normalized records from a device's mDNS response.""" qname = encode_name(SERVICE) @@ -83,7 +102,7 @@ def probe_device(ip, timeout=2.0): records = [] for _ in range(an + ns): name, name_end = read_name(data, off) - rtype, _rclass, ttl, rlen = struct.unpack(">HHIH", data[name_end:name_end + 10]) + rtype, _rclass, ttl, rlen = decode_rr_header(data, name_end) rec = {"name": name, "type": rtype, "ttl": min(ttl, 4500)} if rtype == 12: # PTR rec["ptr_target"] = read_name(data, name_end + 10)[0] @@ -147,7 +166,7 @@ def main(): recs = probe_device(ip) with lock: if recs: - state[ip] = recs + state[ip] = merge_records(state.get(ip, []), recs) log.info("probe %s: %s", ip, f"{len(recs)} records" if recs else "no response") def replay(): diff --git a/netinfra/mdns/test_hap_bridge.py b/netinfra/mdns/test_hap_bridge.py new file mode 100644 index 0000000..30d8b01 --- /dev/null +++ b/netinfra/mdns/test_hap_bridge.py @@ -0,0 +1,62 @@ +#!/usr/bin/env python3 +"""Unit tests for hap-bridge record handling. [#619] + +Run: python3 -m unittest discover -s netinfra/mdns -p 'test_*.py' +""" + +import importlib.util +import os +import unittest + +_SPEC = importlib.util.spec_from_file_location( + "hap_bridge", os.path.join(os.path.dirname(__file__), "hap-bridge.py") +) +hap_bridge = importlib.util.module_from_spec(_SPEC) +_SPEC.loader.exec_module(hap_bridge) + +PTR = {"name": "_hap._tcp.local", "type": 12, "ttl": 4500, "ptr_target": "dev1._hap._tcp.local"} +SRV_OLD = {"name": "dev1._hap._tcp.local", "type": 33, "ttl": 120, "srv_prio": 0, + "srv_weight": 0, "srv_port": 8080, "srv_target": "dev1.local"} +SRV_NEW = {"name": "dev1._hap._tcp.local", "type": 33, "ttl": 120, "srv_prio": 0, + "srv_weight": 0, "srv_port": 8081, "srv_target": "dev1.local"} +TXT = {"name": "dev1._hap._tcp.local", "type": 16, "ttl": 4500, "rdata_hex": "0141"} + + +class TestMergeRecords(unittest.TestCase): + def keyset(self, recs): + return {(r["name"], r["type"]): r for r in recs} + + def test_empty_prev_starts_fresh(self): + merged = hap_bridge.merge_records([], [PTR, SRV_OLD, TXT]) + self.assertEqual(len(merged), 3) + self.assertEqual(self.keyset(merged)[(PTR["name"], 12)], PTR) + + def test_same_key_replaces(self): + merged = hap_bridge.merge_records([PTR, SRV_OLD], [SRV_NEW]) + by_key = self.keyset(merged) + self.assertEqual(by_key[(SRV_NEW["name"], 33)]["srv_port"], 8081) + self.assertEqual(by_key[(PTR["name"], 12)], PTR) + + def test_new_keys_union_in(self): + merged = hap_bridge.merge_records([PTR], [SRV_NEW, TXT]) + self.assertEqual(len(merged), 3) + + def test_empty_new_keeps_prev(self): + merged = hap_bridge.merge_records([PTR, SRV_OLD], []) + self.assertEqual(len(merged), 2) + self.assertEqual(self.keyset(merged)[(SRV_OLD["name"], 33)]["srv_port"], 8080) + + +class TestWireRoundTrip(unittest.TestCase): + def test_ptr_record_encodes_without_compression(self): + wire = hap_bridge.encode_record(PTR) + name, off = hap_bridge.read_name(wire, 0) + self.assertEqual(name, "_hap._tcp.local") + rtype, rclass, ttl, rlen = hap_bridge.decode_rr_header(wire, off) + self.assertEqual((rtype, rclass, rlen), (12, 1, len(hap_bridge.encode_name(PTR["ptr_target"])))) + self.assertEqual(ttl, 4500) + self.assertNotIn(b"\xc0", wire) # no compression pointers in replay + + +if __name__ == "__main__": + unittest.main()