init: zero-dep dns-over-https resolver
This commit is contained in:
commit
7348f9b0dd
18 changed files with 589 additions and 0 deletions
29
tests/test_cache.py
Normal file
29
tests/test_cache.py
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
import time
|
||||
|
||||
from pydoh.cache import TTLCache
|
||||
|
||||
|
||||
def test_cache_hit():
|
||||
cache = TTLCache()
|
||||
cache.set("k", ["1.2.3.4"], ttl=60)
|
||||
assert cache.get("k") == ["1.2.3.4"]
|
||||
|
||||
|
||||
def test_cache_expires():
|
||||
cache = TTLCache()
|
||||
cache.set("k", ["1.2.3.4"], ttl=0.05)
|
||||
assert cache.get("k") == ["1.2.3.4"]
|
||||
time.sleep(0.1)
|
||||
assert cache.get("k") is None
|
||||
|
||||
|
||||
def test_cache_clear():
|
||||
cache = TTLCache()
|
||||
cache.set("k", ["1.2.3.4"], ttl=60)
|
||||
cache.clear()
|
||||
assert cache.get("k") is None
|
||||
|
||||
|
||||
def test_cache_miss():
|
||||
cache = TTLCache()
|
||||
assert cache.get("missing") is None
|
||||
7
tests/test_live.py
Normal file
7
tests/test_live.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
from pydoh import resolve
|
||||
|
||||
|
||||
def test_live_resolve_example_com():
|
||||
ips = resolve("example.com", use_cache=False)
|
||||
assert len(ips) > 0
|
||||
assert all(part.isdigit() for ip in ips for part in ip.split("."))
|
||||
28
tests/test_patch.py
Normal file
28
tests/test_patch.py
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
import socket
|
||||
|
||||
from pydoh.patch import patch_socket, unpatch_socket
|
||||
|
||||
|
||||
def test_patch_and_unpatch_restores_original():
|
||||
original = socket.getaddrinfo
|
||||
patch_socket()
|
||||
assert socket.getaddrinfo is not original
|
||||
unpatch_socket()
|
||||
assert socket.getaddrinfo is original
|
||||
|
||||
|
||||
def test_patch_is_idempotent():
|
||||
patch_socket()
|
||||
patched = socket.getaddrinfo
|
||||
patch_socket()
|
||||
assert socket.getaddrinfo is patched
|
||||
unpatch_socket()
|
||||
|
||||
|
||||
def test_ip_literal_bypasses_doh():
|
||||
patch_socket()
|
||||
try:
|
||||
result = socket.getaddrinfo("127.0.0.1", 80)
|
||||
assert result
|
||||
finally:
|
||||
unpatch_socket()
|
||||
73
tests/test_resolver.py
Normal file
73
tests/test_resolver.py
Normal file
|
|
@ -0,0 +1,73 @@
|
|||
import struct
|
||||
|
||||
import pydoh.resolver as resolver_module
|
||||
from pydoh.resolver import ResolveError, resolve
|
||||
|
||||
|
||||
def _make_response(query_id, ip):
|
||||
header = struct.pack(">HHHHHH", query_id, 0x8180, 1, 1, 0, 0)
|
||||
name = b"\x07example\x03com\x00"
|
||||
question = name + struct.pack(">HH", 1, 1)
|
||||
rdata = bytes(int(part) for part in ip.split("."))
|
||||
answer = b"\xc0\x0c" + struct.pack(">HHIH", 1, 1, 300, 4) + rdata
|
||||
return header + question + answer
|
||||
|
||||
|
||||
def test_resolve_returns_ip(monkeypatch):
|
||||
def fake_doh_request(provider, query, timeout):
|
||||
query_id = struct.unpack(">H", query[:2])[0]
|
||||
return _make_response(query_id, "93.184.216.34")
|
||||
|
||||
monkeypatch.setattr(resolver_module, "_doh_request", fake_doh_request)
|
||||
resolver_module._cache.clear()
|
||||
|
||||
ips = resolve("example.com", use_cache=False)
|
||||
assert ips == ["93.184.216.34"]
|
||||
|
||||
|
||||
def test_resolve_uses_cache(monkeypatch):
|
||||
calls = []
|
||||
|
||||
def fake_doh_request(provider, query, timeout):
|
||||
calls.append(1)
|
||||
query_id = struct.unpack(">H", query[:2])[0]
|
||||
return _make_response(query_id, "1.1.1.1")
|
||||
|
||||
monkeypatch.setattr(resolver_module, "_doh_request", fake_doh_request)
|
||||
resolver_module._cache.clear()
|
||||
|
||||
resolve("example.com", use_cache=True)
|
||||
resolve("example.com", use_cache=True)
|
||||
assert len(calls) == 1
|
||||
|
||||
|
||||
def test_resolve_falls_back_to_next_provider(monkeypatch):
|
||||
attempts = []
|
||||
|
||||
def fake_doh_request(provider, query, timeout):
|
||||
attempts.append(provider.name)
|
||||
if provider.name == "cloudflare":
|
||||
raise RuntimeError("simulated failure")
|
||||
query_id = struct.unpack(">H", query[:2])[0]
|
||||
return _make_response(query_id, "8.8.8.8")
|
||||
|
||||
monkeypatch.setattr(resolver_module, "_doh_request", fake_doh_request)
|
||||
resolver_module._cache.clear()
|
||||
|
||||
ips = resolve("example.com", use_cache=False)
|
||||
assert ips == ["8.8.8.8"]
|
||||
assert attempts[0] == "cloudflare"
|
||||
|
||||
|
||||
def test_resolve_raises_when_all_providers_fail(monkeypatch):
|
||||
def fake_doh_request(provider, query, timeout):
|
||||
raise RuntimeError("down")
|
||||
|
||||
monkeypatch.setattr(resolver_module, "_doh_request", fake_doh_request)
|
||||
resolver_module._cache.clear()
|
||||
|
||||
try:
|
||||
resolve("example.com", use_cache=False)
|
||||
assert False, "expected ResolveError"
|
||||
except ResolveError:
|
||||
pass
|
||||
48
tests/test_wire.py
Normal file
48
tests/test_wire.py
Normal file
|
|
@ -0,0 +1,48 @@
|
|||
import struct
|
||||
|
||||
from pydoh.wire import build_query, parse_response
|
||||
|
||||
|
||||
def test_build_query_shape():
|
||||
query, query_id = build_query("example.com")
|
||||
assert isinstance(query, bytes)
|
||||
assert 0 <= query_id <= 0xFFFF
|
||||
assert query[12:20] == b"\x07example"
|
||||
|
||||
|
||||
def test_parse_response_a_record():
|
||||
query, query_id = build_query("example.com")
|
||||
question = query[12:]
|
||||
|
||||
header = struct.pack(">HHHHHH", query_id, 0x8180, 1, 1, 0, 0)
|
||||
answer = b"\xc0\x0c" + struct.pack(">HHIH", 1, 1, 300, 4) + bytes([93, 184, 216, 34])
|
||||
response = header + question + answer
|
||||
|
||||
resp_id, answers = parse_response(response)
|
||||
assert resp_id == query_id
|
||||
assert answers == [("93.184.216.34", 300)]
|
||||
|
||||
|
||||
def test_parse_response_aaaa_record():
|
||||
query, query_id = build_query("example.com", record_type=28)
|
||||
question = query[12:]
|
||||
|
||||
rdata = bytes.fromhex("20010db8000000000000000000000001")
|
||||
header = struct.pack(">HHHHHH", query_id, 0x8180, 1, 1, 0, 0)
|
||||
answer = b"\xc0\x0c" + struct.pack(">HHIH", 28, 1, 60, 16) + rdata
|
||||
response = header + question + answer
|
||||
|
||||
resp_id, answers = parse_response(response)
|
||||
assert resp_id == query_id
|
||||
assert answers == [("2001:0db8:0000:0000:0000:0000:0000:0001", 60)]
|
||||
|
||||
|
||||
def test_parse_response_no_answers():
|
||||
query, query_id = build_query("nx.example.com")
|
||||
question = query[12:]
|
||||
header = struct.pack(">HHHHHH", query_id, 0x8183, 1, 0, 0, 0)
|
||||
response = header + question
|
||||
|
||||
resp_id, answers = parse_response(response)
|
||||
assert resp_id == query_id
|
||||
assert answers == []
|
||||
Loading…
Add table
Add a link
Reference in a new issue