Files
vpn-loopback-firewall/channel/client.py
T
lou 96d44569fd L2/L3 通道: HTTP/DNS 套壳转发(本地分流代理 + 云端转发器)
- proto.py: 帧协议 [type][req_id][len][payload], req_id 多路复用
- client.py: TLS 长连 + AUTH + 断线自动重连(全标准库, OpenWrt 可跑)
- proxy_http.py: 本地 HTTP 透明代理(明文 80 走通道, 失败回 502)
- proxy_dns.py: DNS 分流(内网域 -> 学校 DNS, 公网域 -> 云通道递归)
- server.py: 云端 AUTH + HTTPS 套壳(统一 UA, 剔 hop-by-hop)+ DNS 递归
- deploy/hook.sh: REDIRECT 规则(内网段放行, 认证 eportal 绝不走通道)
- 30 个新单测(40 总)+ smoke_channel 端到端冒烟 PASS
2026-09-09 14:55:02 +08:00

218 lines
6.9 KiB
Python

"""L2 通道客户端(软路由侧)。
职责: 维护一条到云服务器的 TLS 长连接, 多路复用 HTTP/DNS 请求帧。
- req_id 多路复用: HTTP 大响应不阻塞 DNS 小查询
- AUTH 鉴权(首帧 token), 防公网滥用
- 断线自动重连(最多 retries 次)
- 全标准库(socket/ssl/threading/queue), OpenWrt python3 可直接跑
"""
from __future__ import annotations
import queue
import socket
import ssl
import threading
import time
from typing import Any
from .proto import (
FrameReader,
TYPE_AUTH,
TYPE_AUTH_OK,
TYPE_DNS_REQ,
TYPE_DNS_RESP,
TYPE_ERROR,
TYPE_HTTP_REQ,
TYPE_HTTP_RESP,
TYPE_PING,
TYPE_PONG,
pack_frame,
)
CONNECT_TIMEOUT = 10
REQUEST_TIMEOUT = 90
_ERR = object() # waiter 哨兵: 连接中断
class ChannelError(RuntimeError):
pass
class ChannelClient:
"""到云服务器的通道客户端。线程安全: 任意线程可并发 request。"""
def __init__(
self,
host: str,
port: int,
token: str,
ca_path: str = "",
heartbeat: float = 30.0,
retries: int = 3,
) -> None:
if not host or not token:
raise ValueError("host/token required")
self.host = host
self.port = port
self.token = token
self.ca_path = ca_path
self.heartbeat = heartbeat
self.retries = retries
self._sock: socket.socket | None = None
self._reader_stop = threading.Event()
self._reader_thread: threading.Thread | None = None
self._write_lock = threading.RLock()
self._waiters_lock = threading.Lock()
self._waiters: dict[int, queue.SimpleQueue[Any]] = {}
self._next_id = 1
self._id_lock = threading.Lock()
self._alive = False
self._heartbeat_thread: threading.Thread | None = None
def _alloc_id(self) -> int:
with self._id_lock:
req_id = self._next_id
self._next_id = (self._next_id + 1) & 0xFFFFFFFF
return req_id
def _make_ssl_context(self) -> ssl.SSLContext:
ctx = ssl.create_default_context()
if self.ca_path:
ctx.load_verify_locations(cafile=self.ca_path)
ctx.check_hostname = True
else:
# 自签证书部署: 关闭对端校验, 安全靠 token 鉴权(生产应换 CA 证书)
ctx.check_hostname = False
ctx.verify_mode = ssl.CERT_NONE
return ctx
def _open_socket(self) -> None:
ctx = self._make_ssl_context()
raw = socket.create_connection(
(self.host, self.port), timeout=CONNECT_TIMEOUT
)
try:
sock = ctx.wrap_socket(raw, server_hostname=self.host)
except Exception:
raw.close()
raise
sock.settimeout(None)
self._sock = sock
self._reader_stop.clear()
self._reader_thread = threading.Thread(
target=self._reader_loop, name="channel-reader", daemon=True
)
self._reader_thread.start()
self._alive = True
def _reader_loop(self) -> None:
reader = FrameReader()
assert self._sock is not None
while not self._reader_stop.is_set():
try:
data = self._sock.recv(65536)
except OSError:
break
if not data:
break
reader.feed(data)
while True:
frame = reader.poll()
if frame is None:
break
self._dispatch(*frame)
self._fail_all()
def _dispatch(self, frame_type: int, req_id: int, payload: bytes) -> None:
if frame_type == TYPE_PING:
self._send_raw(pack_frame(TYPE_PONG, req_id))
return
with self._waiters_lock:
q = self._waiters.pop(req_id, None)
if q is not None:
if frame_type == TYPE_ERROR:
q.put(_ERR)
else:
q.put(payload)
def _send_raw(self, data: bytes) -> None:
with self._write_lock:
if self._sock is None:
raise ChannelError("not connected")
self._sock.sendall(data)
def _fail_all(self) -> None:
self._alive = False
with self._waiters_lock:
waiters = list(self._waiters.values())
self._waiters.clear()
for q in waiters:
q.put(_ERR)
def connect(self) -> None:
last_err: Exception | None = None
for attempt in range(self.retries):
try:
self._open_socket()
resp = self.request(TYPE_AUTH, self.token.encode("utf-8"), timeout=10)
if resp is None:
raise ChannelError("auth failed: connection dropped")
return
except Exception as exc: # noqa: BLE001 - 重连必须兜住所有异常
last_err = exc
self.close()
if attempt < self.retries - 1:
time.sleep(1.5 * (attempt + 1))
raise ChannelError(f"connect failed after {self.retries} tries: {last_err}")
def ensure_connected(self) -> None:
"""调用方在 request 抛 ChannelError 后重连用(内部锁防并发)。"""
with self._write_lock:
if self._sock is None or not self._alive:
self.connect()
def request(
self, frame_type: int, payload: bytes = b"", timeout: float = REQUEST_TIMEOUT
) -> bytes | None:
"""发送请求帧并等待响应; 超时/断线返回 None(ERROR 帧也返回 None)。"""
q: queue.SimpleQueue[Any] = queue.SimpleQueue()
req_id = self._alloc_id()
with self._waiters_lock:
self._waiters[req_id] = q
try:
self._send_raw(pack_frame(frame_type, req_id, payload))
except Exception:
with self._waiters_lock:
self._waiters.pop(req_id, None)
if not self._alive:
raise ChannelError("connection lost") from None
raise
try:
result = q.get(timeout=timeout)
except queue.Empty:
with self._waiters_lock:
self._waiters.pop(req_id, None)
return None
if result is _ERR:
raise ChannelError("connection lost")
return result # type: ignore[return-value]
def send_http_request(self, raw_request: bytes) -> bytes | None:
return self.request(TYPE_HTTP_REQ, raw_request)
def send_dns_query(self, raw_query: bytes) -> bytes | None:
return self.request(TYPE_DNS_REQ, raw_query, timeout=15)
def close(self) -> None:
self._reader_stop.set()
self._alive = False
with self._write_lock:
sock = self._sock
self._sock = None
if sock is not None:
try:
sock.close()
except OSError:
pass
self._fail_all()