Files
vpn-loopback-firewall/channel/proxy_dns.py
T

137 lines
4.3 KiB
Python
Raw Normal View History

"""L2 本地 DNS 分流代理(软路由侧)。
接入点: 软路由 smartdns/dnsmasq 的上游指向本代理(127.0.0.1#8053),
全量 DNS 查询先到这里, 按域名分流:
- 内网/学校域名(配置 internal_suffixes, 如 .internal.example)-> 本地转发学校 DNS
(学校内网记录只有学校 DNS 有, 云端解析不到, 认证域名另走 hosts)
- 其余公网域名 -> 走云通道, 云端递归解析(223.5.5.5)
query 全程原样透传(不改 ID), 响应沿原路返回。
"""
from __future__ import annotations
import socket
import threading
from typing import Any
from .client import ChannelClient, ChannelError
from .dnsmsg import parse_qname
DNS_TIMEOUT = 4.0
def _suffix_match(name: str, suffixes: list[str]) -> bool:
for suffix in suffixes:
s = suffix.strip().lower()
if not s:
continue
if s.startswith("."):
if name == s[1:] or name.endswith(s):
return True
elif name == s or name.endswith("." + s):
return True
return False
class DnsSplitProxy:
"""UDP 53 分流代理。"""
def __init__(
self,
client: ChannelClient,
listen_port: int,
school_dns: str,
internal_suffixes: list[str],
) -> None:
if listen_port < 1 or listen_port > 65535:
raise ValueError(f"bad listen port: {listen_port}")
self.client = client
self.port = listen_port
self.school_dns = school_dns
self.internal_suffixes = internal_suffixes
self.local_resolved = 0
self.channel_resolved = 0
self.failed = 0
def serve(self) -> None:
srv = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
srv.bind(("127.0.0.1", self.port))
print(
f"[dns-proxy] listen 127.0.0.1:{self.port} "
f"school_dns={self.school_dns}",
flush=True,
)
while True:
data, client_addr = srv.recvfrom(4096)
threading.Thread(
target=self._handle,
args=(srv, data, client_addr),
name="dns-proxy-query",
daemon=True,
).start()
def _handle(
self,
srv: socket.socket,
query: bytes,
client_addr: tuple[str, int],
) -> None:
name = parse_qname(query)
try:
if name is not None and _suffix_match(name, self.internal_suffixes):
self._local_forward(srv, query, client_addr)
self.local_resolved += 1
else:
self._channel_forward(srv, query, client_addr)
self.channel_resolved += 1
except OSError:
self.failed += 1
def _local_forward(
self, srv: socket.socket, query: bytes, client_addr: tuple[str, int]
) -> None:
"""直连学校 DNS(本地网络, 不走通道)。"""
up = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
try:
up.settimeout(DNS_TIMEOUT)
up.sendto(query, (self.school_dns, 53))
resp, _ = up.recvfrom(4096)
srv.sendto(resp, client_addr)
finally:
up.close()
def _channel_forward(
self, srv: socket.socket, query: bytes, client_addr: tuple[str, int]
) -> None:
resp = self._channel_send(query)
if resp is not None:
srv.sendto(resp, client_addr)
def _channel_send(self, query: bytes) -> bytes | None:
try:
return self.client.send_dns_query(query)
except ChannelError:
try:
self.client.ensure_connected()
return self.client.send_dns_query(query)
except ChannelError:
return None
def run_dns_proxy(cfg: dict[str, Any]) -> None:
cloud: dict[str, Any] = cfg["cloud"]
proxy_cfg: dict[str, Any] = cfg["dns_proxy"]
client = ChannelClient(
host=str(cloud["host"]),
port=int(cloud["port"]),
token=str(cloud["token"]),
ca_path=str(cloud.get("ca_path", "")),
heartbeat=float(cfg.get("heartbeat", 30)),
)
client.connect()
DnsSplitProxy(
client,
listen_port=int(proxy_cfg["listen_port"]),
school_dns=str(proxy_cfg["school_dns"]),
internal_suffixes=[str(s) for s in proxy_cfg.get("internal_suffixes", [])],
).serve()