fix(mdns): accumulate probe records across responses; unit tests [#619]
Live proof: ecobee answers unicast mDNS with PTR only (_hap._tcp.local -> "Main Floor._hap._tcp.local"), so the old replace-on-probe wiped learned records every cycle. Merge by (name, type) instead. CONFIG_PATH now env-overridable for tests. Details: https://projects.knownelement.com/issues/619#note-5 💘 Generated with Crush Assisted-by: Crush:glm-5.2
This commit is contained in:
@@ -19,13 +19,14 @@ Design notes:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
|
import os
|
||||||
import logging
|
import logging
|
||||||
import socket
|
import socket
|
||||||
import struct
|
import struct
|
||||||
import threading
|
import threading
|
||||||
import time
|
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_GROUP = "224.0.0.251"
|
||||||
MDNS_PORT = 5353
|
MDNS_PORT = 5353
|
||||||
REFRESH_INTERVAL = 60
|
REFRESH_INTERVAL = 60
|
||||||
@@ -60,6 +61,24 @@ def encode_name(name):
|
|||||||
return out + b"\x00"
|
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):
|
def probe_device(ip, timeout=2.0):
|
||||||
"""Return normalized records from a device's mDNS response."""
|
"""Return normalized records from a device's mDNS response."""
|
||||||
qname = encode_name(SERVICE)
|
qname = encode_name(SERVICE)
|
||||||
@@ -83,7 +102,7 @@ def probe_device(ip, timeout=2.0):
|
|||||||
records = []
|
records = []
|
||||||
for _ in range(an + ns):
|
for _ in range(an + ns):
|
||||||
name, name_end = read_name(data, off)
|
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)}
|
rec = {"name": name, "type": rtype, "ttl": min(ttl, 4500)}
|
||||||
if rtype == 12: # PTR
|
if rtype == 12: # PTR
|
||||||
rec["ptr_target"] = read_name(data, name_end + 10)[0]
|
rec["ptr_target"] = read_name(data, name_end + 10)[0]
|
||||||
@@ -147,7 +166,7 @@ def main():
|
|||||||
recs = probe_device(ip)
|
recs = probe_device(ip)
|
||||||
with lock:
|
with lock:
|
||||||
if recs:
|
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")
|
log.info("probe %s: %s", ip, f"{len(recs)} records" if recs else "no response")
|
||||||
|
|
||||||
def replay():
|
def replay():
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user