Skip to content

Commit b0a399c

Browse files
committed
Add DNS query test script
1 parent 2833dd3 commit b0a399c

2 files changed

Lines changed: 164 additions & 0 deletions

File tree

.gitignore

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
__pycache__
12
dns2socks.xcframework/
23
Cargo.lock
34
.vscode/

scripts/test_dns_query.py

Lines changed: 163 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,163 @@
1+
#!/usr/bin/env python3
2+
"""Send a DNS query for baidu.com to a local resolver.
3+
4+
Defaults:
5+
- host: 127.0.0.1
6+
- port: 53
7+
- name: baidu.com
8+
- type: A
9+
- transport: TCP
10+
11+
Examples:
12+
python3 scripts/test_dns_query.py
13+
python3 scripts/test_dns_query.py --udp
14+
python3 scripts/test_dns_query.py --name baidu.com --server 127.0.0.1 --port 53
15+
"""
16+
17+
from __future__ import annotations
18+
19+
import argparse
20+
import random
21+
import socket
22+
import struct
23+
import sys
24+
from typing import List, Tuple
25+
26+
27+
def encode_name(name: str) -> bytes:
28+
labels = name.rstrip(".").split(".")
29+
encoded = bytearray()
30+
for label in labels:
31+
label_bytes = label.encode("ascii")
32+
if len(label_bytes) > 63:
33+
raise ValueError(f"label too long: {label}")
34+
encoded.append(len(label_bytes))
35+
encoded.extend(label_bytes)
36+
encoded.append(0)
37+
return bytes(encoded)
38+
39+
40+
def decode_name(message: bytes, offset: int) -> Tuple[str, int]:
41+
labels: List[str] = []
42+
jumped = False
43+
original_offset = offset
44+
45+
while True:
46+
length = message[offset]
47+
if length == 0:
48+
offset += 1
49+
break
50+
51+
if length & 0xC0 == 0xC0:
52+
pointer = ((length & 0x3F) << 8) | message[offset + 1]
53+
if not jumped:
54+
original_offset = offset + 2
55+
offset = pointer
56+
jumped = True
57+
continue
58+
59+
offset += 1
60+
labels.append(message[offset : offset + length].decode("ascii", errors="replace"))
61+
offset += length
62+
63+
return ".".join(labels), (original_offset if jumped else offset)
64+
65+
66+
def build_query(name: str, qtype: int = 1) -> Tuple[int, bytes]:
67+
query_id = random.randint(0, 0xFFFF)
68+
flags = 0x0100 # recursion desired
69+
header = struct.pack("!HHHHHH", query_id, flags, 1, 0, 0, 0)
70+
question = encode_name(name) + struct.pack("!HH", qtype, 1)
71+
return query_id, header + question
72+
73+
74+
def parse_response(message: bytes, expected_id: int) -> None:
75+
if len(message) < 12:
76+
raise ValueError("response too short")
77+
78+
(response_id, flags, qdcount, ancount, nscount, arcount) = struct.unpack("!HHHHHH", message[:12])
79+
if response_id != expected_id:
80+
raise ValueError(f"unexpected response id: {response_id} != {expected_id}")
81+
82+
rcode = flags & 0x000F
83+
if rcode != 0:
84+
raise RuntimeError(f"dns error rcode={rcode}")
85+
86+
offset = 12
87+
for _ in range(qdcount):
88+
_, offset = decode_name(message, offset)
89+
offset += 4
90+
91+
answers = []
92+
for _ in range(ancount):
93+
name, offset = decode_name(message, offset)
94+
rtype, rclass, ttl, rdlength = struct.unpack("!HHIH", message[offset : offset + 10])
95+
offset += 10
96+
rdata = message[offset : offset + rdlength]
97+
offset += rdlength
98+
99+
if rtype == 1 and rclass == 1 and rdlength == 4:
100+
ip = socket.inet_ntoa(rdata)
101+
answers.append((name, "A", ttl, ip))
102+
elif rtype == 28 and rclass == 1 and rdlength == 16:
103+
ip = socket.inet_ntop(socket.AF_INET6, rdata)
104+
answers.append((name, "AAAA", ttl, ip))
105+
else:
106+
answers.append((name, f"TYPE{rtype}", ttl, rdata.hex()))
107+
108+
print(f"answers={ancount} ns={nscount} additional={arcount}")
109+
for name, record_type, ttl, value in answers:
110+
print(f"{name} {ttl} IN {record_type} {value}")
111+
112+
113+
def send_udp(server: str, port: int, payload: bytes, timeout: float) -> bytes:
114+
with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as sock:
115+
sock.settimeout(timeout)
116+
sock.sendto(payload, (server, port))
117+
data, _ = sock.recvfrom(4096)
118+
return data
119+
120+
121+
def send_tcp(server: str, port: int, payload: bytes, timeout: float) -> bytes:
122+
with socket.create_connection((server, port), timeout=timeout) as sock:
123+
sock.settimeout(timeout)
124+
sock.sendall(struct.pack("!H", len(payload)) + payload)
125+
header = sock.recv(2)
126+
if len(header) < 2:
127+
raise RuntimeError("short TCP DNS header")
128+
(length,) = struct.unpack("!H", header)
129+
data = bytearray()
130+
while len(data) < length:
131+
chunk = sock.recv(length - len(data))
132+
if not chunk:
133+
break
134+
data.extend(chunk)
135+
if len(data) != length:
136+
raise RuntimeError("short TCP DNS body")
137+
return bytes(data)
138+
139+
140+
def main() -> int:
141+
parser = argparse.ArgumentParser(description="Query baidu.com through a local DNS server")
142+
parser.add_argument("--server", default="127.0.0.1")
143+
parser.add_argument("--port", type=int, default=53)
144+
parser.add_argument("--name", default="baidu.com")
145+
parser.add_argument("--timeout", type=float, default=5.0)
146+
parser.add_argument("--udp", action="store_true", help="use UDP instead of TCP")
147+
args = parser.parse_args()
148+
149+
query_id, payload = build_query(args.name)
150+
transport = "UDP" if args.udp else "TCP"
151+
print(f"query={args.name} server={args.server}:{args.port} transport={transport}")
152+
153+
if args.udp:
154+
response = send_udp(args.server, args.port, payload, args.timeout)
155+
else:
156+
response = send_tcp(args.server, args.port, payload, args.timeout)
157+
158+
parse_response(response, query_id)
159+
return 0
160+
161+
162+
if __name__ == "__main__":
163+
sys.exit(main())

0 commit comments

Comments
 (0)