Add umedbazarov/ruh-vpn: VPN/proxy manager for sing-box (#304)
* Add umedbazarov/ruh-vpn: VPN/proxy manager for sing-box New community plugin: bar widget, panel, service and control-center shortcut for managing SSH, VLESS, VMess, Shadowsocks and SOCKS5 connections through sing-box, with routing presets, custom rules, system-proxy/TUN modes and a kill switch. The bundled Python backend serves a loopback control API protected by a per-launch bearer token. * Address review: sanitize kill-switch ruleset, scope TUN capability, fix mux error path, disclose DNS - kill switch: only pre-resolved, canonicalized literal IPs enter the nft ruleset; domains are resolved first and anything unparseable is dropped, so subscription-supplied addresses can no longer inject nft syntax - TUN: CAP_NET_ADMIN is granted to a plugin-private copy of sing-box in a 0700 directory instead of the shared system binary; the copy is refreshed (clearing the cap) when the system binary changes, and the legacy grant on the shared binary is removed in the same polkit prompt - fix NameError in the mux startup failure path (undefined mux_name) that hid the log tail and skipped teardown - README: disclose plain-UDP DNS endpoints (8.8.8.8 via tunnel, 223.5.5.5 direct in rules mode) alongside the TUN DoH endpoint --------- Co-authored-by: Umedjon Bazarov <170195993+UmedjonBA@users.noreply.github.com>
This commit is contained in:
@@ -0,0 +1,202 @@
|
||||
"""Entry point: start asyncio loop, bootstrap service, expose HTTP control API.
|
||||
|
||||
The Luau [[service]] entry (service.luau) launches this module with
|
||||
noctalia.runStream() and consumes the newline-JSON events we print to stdout.
|
||||
Commands come back in over the loopback HTTP control port.
|
||||
|
||||
Port selection: RUH_VPN_CONTROL_PORT env var, else 11090.
|
||||
|
||||
Coexistence safety: if the control port is already bound, another backend is
|
||||
already running (e.g. the user's active proxy). We emit a "port-in-use" error
|
||||
and exit WITHOUT running the destructive shutdown path, so we never tear down a
|
||||
proxy this instance does not own. The Luau service probes /healthz first and
|
||||
only spawns us when no backend is present, so this is a belt-and-suspenders
|
||||
guard.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import secrets
|
||||
import signal
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
from backend.http.control import DEFAULT_PORT, HOST, emit, serve
|
||||
from backend.identity import PREFIX
|
||||
from backend.paths import RUNTIME_DIR, ensure_private_dir, protect_file
|
||||
from backend.service.vpn_service import VpnService
|
||||
|
||||
PIDFILE = RUNTIME_DIR / f"{PREFIX}-backend.pid"
|
||||
TOKENFILE = RUNTIME_DIR / f"{PREFIX}-control.token"
|
||||
DETACHED_LOG = RUNTIME_DIR / f"{PREFIX}-backend-detached.log"
|
||||
|
||||
|
||||
class _StdoutGuard:
|
||||
"""A stdout that outlives its reader.
|
||||
|
||||
service.luau spawns us through noctalia.runStream(), so stdout is a pipe
|
||||
owned by that shell. When the shell exits we are meant to keep running (the
|
||||
next shell re-attaches over /healthz and the proxy survives), but the read
|
||||
end of the pipe closes with it, and from then on every write raises
|
||||
BrokenPipeError. Because _now_log() prints, that exception surfaced inside
|
||||
whichever RPC logged first: StartProxy died on its very first log line and
|
||||
the proxy silently never started. A vanished reader must not be fatal, so
|
||||
fall back to a file and carry on.
|
||||
"""
|
||||
|
||||
def __init__(self, stream: object, fallback: Path) -> None:
|
||||
self._stream = stream
|
||||
self._fallback = fallback
|
||||
self._demoted = False
|
||||
|
||||
def _demote(self) -> None:
|
||||
# Close explicitly: the dead pipe still holds whatever we buffered, and
|
||||
# letting the finalizer discover that prints "Exception ignored".
|
||||
old, self._stream = self._stream, None
|
||||
if old is not None:
|
||||
try:
|
||||
old.close()
|
||||
except Exception:
|
||||
pass
|
||||
if self._demoted:
|
||||
return
|
||||
self._demoted = True
|
||||
try:
|
||||
self._stream = open(self._fallback, "a", buffering=1)
|
||||
protect_file(self._fallback)
|
||||
except OSError:
|
||||
self._stream = None
|
||||
|
||||
def write(self, data: str) -> int:
|
||||
# At most two tries: pipe -> fallback file -> give up silently.
|
||||
for _ in range(2):
|
||||
stream = self._stream
|
||||
if stream is None:
|
||||
break
|
||||
try:
|
||||
written = stream.write(data)
|
||||
# Flush here, not in flush(): the stream is buffered, so a write
|
||||
# to a dead pipe succeeds and only the flush raises. By then the
|
||||
# line is stranded in a buffer we are about to drop. Flushing
|
||||
# while `data` is still in hand lets the retry re-send it.
|
||||
stream.flush()
|
||||
return written
|
||||
except (BrokenPipeError, OSError, ValueError):
|
||||
self._demote()
|
||||
return len(data)
|
||||
|
||||
def flush(self) -> None:
|
||||
# write() already flushed; this only has to stay well-behaved.
|
||||
stream = self._stream
|
||||
if stream is None:
|
||||
return
|
||||
try:
|
||||
stream.flush()
|
||||
except (BrokenPipeError, OSError, ValueError):
|
||||
self._demote()
|
||||
|
||||
def isatty(self) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _install_stdio_guards() -> None:
|
||||
sys.stdout = _StdoutGuard(sys.stdout, DETACHED_LOG)
|
||||
sys.stderr = _StdoutGuard(sys.stderr, DETACHED_LOG)
|
||||
|
||||
|
||||
def _resolve_port() -> int:
|
||||
raw = os.environ.get("RUH_VPN_CONTROL_PORT", "")
|
||||
try:
|
||||
return int(raw) if raw else DEFAULT_PORT
|
||||
except ValueError:
|
||||
return DEFAULT_PORT
|
||||
|
||||
|
||||
def _write_pidfile(port: int) -> None:
|
||||
try:
|
||||
PIDFILE.write_text(f"{os.getpid()} {port}\n")
|
||||
protect_file(PIDFILE)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _remove_pidfile() -> None:
|
||||
try:
|
||||
# Only remove if it still points at us
|
||||
content = PIDFILE.read_text().split()
|
||||
if content and content[0] == str(os.getpid()):
|
||||
PIDFILE.unlink()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _remove_tokenfile(token: str) -> None:
|
||||
try:
|
||||
# Only remove if it still holds our token
|
||||
if TOKENFILE.read_text().strip() == token:
|
||||
TOKENFILE.unlink()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def main() -> int:
|
||||
ensure_private_dir(RUNTIME_DIR)
|
||||
_install_stdio_guards()
|
||||
port = _resolve_port()
|
||||
|
||||
svc = VpnService()
|
||||
await svc.bootstrap()
|
||||
|
||||
# Per-launch RPC token. serve() writes it to TOKENFILE (0600) only after
|
||||
# the port is bound; the Luau service reads the file and sends it as
|
||||
# "Authorization: Bearer <token>" on every /rpc call.
|
||||
token = secrets.token_urlsafe(32)
|
||||
try:
|
||||
control = await serve(svc, port, token=token, token_file=TOKENFILE)
|
||||
except OSError as exc:
|
||||
# Port already in use: another backend owns the proxy. Do NOT shut down.
|
||||
emit({"event": "error", "data": {"message": f"control port {port} in use: {exc}"}})
|
||||
print(f"[error] control port {port} already in use; exiting without teardown", flush=True)
|
||||
return 3
|
||||
|
||||
_write_pidfile(port)
|
||||
print(f"[info] Ruh VPN backend ready on http://{HOST}:{port} (pid {os.getpid()})", flush=True)
|
||||
|
||||
stop_event = asyncio.Event()
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
def _shutdown(*_: object) -> None:
|
||||
if not stop_event.is_set():
|
||||
print("[info] shutdown signal received", flush=True)
|
||||
stop_event.set()
|
||||
|
||||
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||
try:
|
||||
loop.add_signal_handler(sig, _shutdown)
|
||||
except NotImplementedError:
|
||||
pass
|
||||
|
||||
await stop_event.wait()
|
||||
|
||||
print("[info] shutting down VPN service", flush=True)
|
||||
try:
|
||||
await control.stop()
|
||||
except Exception as exc:
|
||||
print(f"[error] control stop error: {exc}", flush=True)
|
||||
try:
|
||||
await svc.shutdown()
|
||||
except Exception as exc:
|
||||
print(f"[error] shutdown error: {exc}", flush=True)
|
||||
_remove_pidfile()
|
||||
_remove_tokenfile(token)
|
||||
emit({"event": "exit", "data": {"code": 0}})
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
sys.exit(asyncio.run(main()))
|
||||
except KeyboardInterrupt:
|
||||
sys.exit(0)
|
||||
@@ -0,0 +1,37 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import aiofiles
|
||||
|
||||
from backend.models.server import Settings
|
||||
from backend.paths import DATA_DIR, ensure_private_dir, protect_file
|
||||
|
||||
SETTINGS_FILE = DATA_DIR / "settings.json"
|
||||
|
||||
|
||||
def ensure_dirs() -> None:
|
||||
ensure_private_dir(DATA_DIR)
|
||||
|
||||
|
||||
async def load_settings() -> Settings:
|
||||
ensure_dirs()
|
||||
if not SETTINGS_FILE.exists():
|
||||
return Settings()
|
||||
try:
|
||||
async with aiofiles.open(SETTINGS_FILE, "r") as f:
|
||||
raw = await f.read()
|
||||
data = json.loads(raw or "{}")
|
||||
return Settings.model_validate(data)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return Settings()
|
||||
|
||||
|
||||
async def save_settings(settings: Settings) -> None:
|
||||
ensure_dirs()
|
||||
tmp = SETTINGS_FILE.with_suffix(".json.tmp")
|
||||
async with aiofiles.open(tmp, "w") as f:
|
||||
await f.write(json.dumps(settings.model_dump(exclude_none=True), indent=2))
|
||||
protect_file(tmp)
|
||||
os.replace(tmp, SETTINGS_FILE)
|
||||
@@ -0,0 +1,55 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections import deque
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Optional
|
||||
|
||||
from backend.models.server import RoutingRule, Server, Settings, StatusInfo
|
||||
|
||||
|
||||
LogEntry = tuple[float, str, str] # (timestamp, level, message)
|
||||
|
||||
|
||||
@dataclass
|
||||
class AppState:
|
||||
servers: list[Server] = field(default_factory=list)
|
||||
rules: list[RoutingRule] = field(default_factory=list)
|
||||
settings: Settings = field(default_factory=Settings)
|
||||
status: StatusInfo = field(default_factory=StatusInfo)
|
||||
pids: dict[str, int] = field(default_factory=dict)
|
||||
logs: deque[LogEntry] = field(default_factory=lambda: deque(maxlen=500))
|
||||
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
|
||||
|
||||
status_listeners: list[Callable[[StatusInfo], None]] = field(default_factory=list)
|
||||
server_list_listeners: list[Callable[[], None]] = field(default_factory=list)
|
||||
log_listeners: list[Callable[[str, str], None]] = field(default_factory=list)
|
||||
|
||||
def get_server(self, server_id: str) -> Optional[Server]:
|
||||
for s in self.servers:
|
||||
if s.id == server_id:
|
||||
return s
|
||||
return None
|
||||
|
||||
def emit_status(self) -> None:
|
||||
for cb in list(self.status_listeners):
|
||||
try:
|
||||
cb(self.status)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def emit_server_list(self) -> None:
|
||||
for cb in list(self.server_list_listeners):
|
||||
try:
|
||||
cb()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def emit_log(self, level: str, message: str) -> None:
|
||||
import time
|
||||
self.logs.append((time.time(), level, message))
|
||||
for cb in list(self.log_listeners):
|
||||
try:
|
||||
cb(level, message)
|
||||
except Exception:
|
||||
pass
|
||||
@@ -0,0 +1,82 @@
|
||||
"""Resolve a server's country code, so the UI can show a flag.
|
||||
|
||||
Nothing else in the backend knows a server's country: the models simply accept
|
||||
the extra key. The lookup is best-effort and always optional — a server with no
|
||||
country just shows no flag, exactly as before.
|
||||
|
||||
Privacy: this asks a third party (api.country.is) "which country is this IP in",
|
||||
which discloses the address of the user's own VPN server to that service. Hence
|
||||
the `geoip_country` plugin setting, which service.luau forwards as
|
||||
RUH_VPN_GEOIP so it can be turned off. The endpoint is HTTPS and returns
|
||||
only {"ip": ..., "country": ...}; an offline answer isn't possible here — no
|
||||
GeoIP database is installed (no *.mmdb) and sing-box's .srs rulesets only cover
|
||||
specific countries (cn/ir).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import ipaddress
|
||||
import os
|
||||
import re
|
||||
|
||||
# Forwarded by service.luau from the geoip_country setting.
|
||||
ENABLED = os.environ.get("RUH_VPN_GEOIP", "1").lower() not in ("0", "false", "no")
|
||||
|
||||
try:
|
||||
import aiohttp
|
||||
except ImportError: # pragma: no cover - matches monitoring/health.py's guard
|
||||
aiohttp = None # type: ignore
|
||||
|
||||
LOOKUP_URL = "https://api.country.is/{ip}"
|
||||
TIMEOUT_SEC = 6
|
||||
|
||||
_CC_RE = re.compile(r"^[A-Za-z]{2}$")
|
||||
|
||||
|
||||
async def _resolve_ip(host: str) -> str | None:
|
||||
"""Return `host` if it is already an IP, else its first A/AAAA record."""
|
||||
try:
|
||||
ipaddress.ip_address(host)
|
||||
return host
|
||||
except ValueError:
|
||||
pass
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
infos = await asyncio.wait_for(
|
||||
loop.getaddrinfo(host, None), timeout=TIMEOUT_SEC
|
||||
)
|
||||
except (OSError, asyncio.TimeoutError):
|
||||
return None
|
||||
return infos[0][4][0] if infos else None
|
||||
|
||||
|
||||
async def lookup_country(host: str) -> str | None:
|
||||
"""Best-effort ISO-3166 alpha-2 (lowercase) for `host`. None on any failure.
|
||||
|
||||
Never raises: a missing flag must not be able to fail an AddServer.
|
||||
"""
|
||||
if not host or aiohttp is None:
|
||||
return None
|
||||
ip = await _resolve_ip(host.strip())
|
||||
if not ip:
|
||||
return None
|
||||
# A private address has no country, and asking would leak nothing useful.
|
||||
try:
|
||||
if not ipaddress.ip_address(ip).is_global:
|
||||
return None
|
||||
except ValueError:
|
||||
return None
|
||||
try:
|
||||
timeout = aiohttp.ClientTimeout(total=TIMEOUT_SEC)
|
||||
async with aiohttp.ClientSession(timeout=timeout) as session:
|
||||
async with session.get(LOOKUP_URL.format(ip=ip)) as resp:
|
||||
if resp.status != 200:
|
||||
return None
|
||||
data = await resp.json(content_type=None)
|
||||
except Exception:
|
||||
return None
|
||||
cc = (data or {}).get("country") if isinstance(data, dict) else None
|
||||
if isinstance(cc, str) and _CC_RE.match(cc):
|
||||
return cc.lower()
|
||||
return None
|
||||
@@ -0,0 +1,224 @@
|
||||
"""Localhost HTTP control interface for the VPN backend.
|
||||
|
||||
Replaces the old DBus interface (backend/dbus/dbus_server.py) for the Luau
|
||||
plugin. The Luau [[service]] entry talks to us two ways:
|
||||
|
||||
* commands → POST http://127.0.0.1:<port>/rpc with body
|
||||
{"method": "StartProxy", "args": [...]}
|
||||
reply {"result": ...} or {"error": "..."}
|
||||
|
||||
* events → we print newline-delimited JSON to *stdout*, which the Luau
|
||||
service consumes via noctalia.runStream():
|
||||
{"event": "StatusChanged", "data": {...}}
|
||||
{"event": "ServerListChanged"}
|
||||
{"event": "LogMessage", "data": {"level","message"}}
|
||||
{"event": "TrafficUpdate", "data": {...}}
|
||||
{"event": "ready", "data": {"port": <port>}}
|
||||
|
||||
Only 127.0.0.1 is bound, and /rpc additionally requires a per-launch bearer
|
||||
token: loopback alone would let any local user (not just the owner) drive the
|
||||
VPN and read server credentials. The token is generated at startup and written
|
||||
to a 0600 file inside the private runtime dir, so only the owning user's
|
||||
processes — the Luau service among them — can read it. /healthz stays open; it
|
||||
carries nothing but liveness and the port, and the Luau service probes it
|
||||
before it knows the token.
|
||||
|
||||
The method map mirrors the DBus contract 1:1. Over JSON, dict arguments arrive
|
||||
as plain Python dicts, so none of the DBus Variant coercion is needed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hmac
|
||||
import inspect
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
from backend.service.vpn_service import VpnService
|
||||
|
||||
HOST = "127.0.0.1"
|
||||
DEFAULT_PORT = 11090
|
||||
|
||||
|
||||
def emit(obj: dict) -> None:
|
||||
"""Write one JSON event line to stdout for the Luau service to stream."""
|
||||
try:
|
||||
sys.stdout.write(json.dumps(obj, ensure_ascii=False, default=str) + "\n")
|
||||
sys.stdout.flush()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
# ------------------------------------------------------------------ dispatch
|
||||
|
||||
# Each handler receives (svc, args) and returns either a value or an awaitable.
|
||||
# Names and argument order match backend/dbus/dbus_server.py exactly.
|
||||
def _build_handlers() -> dict[str, Callable[[VpnService, list], Any]]:
|
||||
return {
|
||||
# ---- lifecycle / status ----
|
||||
"StartProxy": lambda s, a: s.start_proxy(a[0], a[1], a[2]),
|
||||
"StopProxy": lambda s, a: s.stop_proxy(),
|
||||
"GetStatus": lambda s, a: s.get_status(),
|
||||
"GetHealth": lambda s, a: s.get_health(),
|
||||
"GetTrafficStats": lambda s, a: s.get_traffic_stats(),
|
||||
"CheckDnsLeak": lambda s, a: s.check_dns_leak(),
|
||||
"RunSpeedTest": lambda s, a: s.run_speed_test(),
|
||||
# ---- servers ----
|
||||
"GetServers": lambda s, a: s.list_servers(),
|
||||
"AddServer": lambda s, a: s.add_server(a[0]),
|
||||
"UpdateServer": lambda s, a: s.update_server(a[0]),
|
||||
"RemoveServer": lambda s, a: s.remove_server(a[0]),
|
||||
"SwitchServer": lambda s, a: s.switch_server(a[0]),
|
||||
"PingServer": lambda s, a: s.ping(a[0]),
|
||||
"ParseShareLink": lambda s, a: s.add_from_link(a[0]),
|
||||
# ---- modes ----
|
||||
"SetMode": lambda s, a: s.set_mode(a[0]),
|
||||
"SetProxyMode": lambda s, a: s.set_proxy_mode(a[0]),
|
||||
# ---- routing rules ----
|
||||
"GetRoutingRules": lambda s, a: s.list_rules(),
|
||||
"AddRoutingRule": lambda s, a: s.add_rule(a[0]),
|
||||
"RemoveRoutingRule": lambda s, a: s.remove_rule(a[0]),
|
||||
"GetPresets": lambda s, a: s.list_presets(),
|
||||
"TogglePreset": lambda s, a: s.toggle_preset(a[0], a[1]),
|
||||
# ---- kill switch ----
|
||||
"SetKillSwitch": lambda s, a: s.set_kill_switch(a[0]),
|
||||
"GetKillSwitchStatus": lambda s, a: s.get_kill_switch_status(),
|
||||
# ---- subscriptions ----
|
||||
"AddSubscription": lambda s, a: s.add_subscription(a[0], a[1] if len(a) > 1 else ""),
|
||||
"RemoveSubscription": lambda s, a: s.remove_subscription(a[0]),
|
||||
"UpdateSubscription": lambda s, a: s.update_subscription(a[0]),
|
||||
"GetSubscriptions": lambda s, a: s.list_subscriptions(),
|
||||
# ---- logs / settings ----
|
||||
"GetLogs": lambda s, a: s.get_logs(),
|
||||
"GetSettings": lambda s, a: s.get_settings(),
|
||||
"UpdateSettings": lambda s, a: s.update_settings(a[0]),
|
||||
}
|
||||
|
||||
|
||||
class ControlServer:
|
||||
def __init__(
|
||||
self,
|
||||
service: VpnService,
|
||||
port: int = DEFAULT_PORT,
|
||||
token: str = "",
|
||||
token_file: Path | None = None,
|
||||
) -> None:
|
||||
self._svc = service
|
||||
self._port = port
|
||||
self._token = token
|
||||
self._token_file = token_file
|
||||
self._handlers = _build_handlers()
|
||||
self._runner: web.AppRunner | None = None
|
||||
self._wire_events()
|
||||
|
||||
# ------------------------------------------------------------- events
|
||||
def _wire_events(self) -> None:
|
||||
svc = self._svc
|
||||
svc.state.status_listeners.append(self._on_status)
|
||||
svc.state.server_list_listeners.append(self._on_server_list)
|
||||
svc.state.log_listeners.append(self._on_log)
|
||||
svc.add_traffic_listener(self._on_traffic)
|
||||
|
||||
def _on_status(self, status_obj) -> None:
|
||||
try:
|
||||
data = status_obj.model_dump(exclude_none=True)
|
||||
except Exception:
|
||||
return
|
||||
emit({"event": "StatusChanged", "data": data})
|
||||
|
||||
def _on_server_list(self) -> None:
|
||||
emit({"event": "ServerListChanged"})
|
||||
|
||||
def _on_log(self, level: str, message: str) -> None:
|
||||
msg = message if len(message) <= 1024 else message[:1024] + "..."
|
||||
emit({"event": "LogMessage", "data": {"level": level, "message": msg}})
|
||||
|
||||
def _on_traffic(self, stats: dict) -> None:
|
||||
emit({"event": "TrafficUpdate", "data": stats})
|
||||
|
||||
# ------------------------------------------------------------- http
|
||||
def _authorized(self, request: web.Request) -> bool:
|
||||
if not self._token:
|
||||
return False
|
||||
header = request.headers.get("Authorization", "")
|
||||
scheme, _, presented = header.partition(" ")
|
||||
if scheme.lower() != "bearer":
|
||||
return False
|
||||
return hmac.compare_digest(presented.strip(), self._token)
|
||||
|
||||
async def _handle_rpc(self, request: web.Request) -> web.Response:
|
||||
if not self._authorized(request):
|
||||
return web.json_response({"error": "unauthorized"}, status=401)
|
||||
try:
|
||||
req = await request.json()
|
||||
except Exception:
|
||||
return web.json_response({"error": "invalid JSON body"}, status=400)
|
||||
|
||||
method = req.get("method", "")
|
||||
args = req.get("args") or []
|
||||
handler = self._handlers.get(method)
|
||||
if handler is None:
|
||||
return web.json_response({"error": f"unknown method: {method}"}, status=404)
|
||||
|
||||
try:
|
||||
result = handler(self._svc, args)
|
||||
if inspect.isawaitable(result):
|
||||
result = await result
|
||||
except ValueError as exc:
|
||||
return web.json_response({"error": str(exc) or "invalid argument"}, status=400)
|
||||
except IndexError:
|
||||
return web.json_response({"error": f"missing arguments for {method}"}, status=400)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return web.json_response({"error": f"{type(exc).__name__}: {exc}"}, status=500)
|
||||
|
||||
return web.json_response({"result": result}, dumps=lambda o: json.dumps(o, default=str))
|
||||
|
||||
async def _handle_health(self, request: web.Request) -> web.Response:
|
||||
return web.json_response({"ok": True, "port": self._port})
|
||||
|
||||
def _publish_token(self) -> None:
|
||||
"""Write the token file, readable by the owning user only.
|
||||
|
||||
Must run only after the port is bound: a second backend losing the
|
||||
EADDRINUSE race exits without ever binding, and writing earlier would
|
||||
let that loser clobber the live backend's token on its way out."""
|
||||
if self._token_file is None:
|
||||
return
|
||||
fd = os.open(self._token_file, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
|
||||
with os.fdopen(fd, "w") as fh:
|
||||
fh.write(self._token + "\n")
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Bind the control socket. Raises OSError if the port is already taken
|
||||
(another backend is running) — the caller must NOT fall through to the
|
||||
destructive shutdown path in that case."""
|
||||
app = web.Application()
|
||||
app.router.add_post("/rpc", self._handle_rpc)
|
||||
app.router.add_get("/healthz", self._handle_health)
|
||||
self._runner = web.AppRunner(app)
|
||||
await self._runner.setup()
|
||||
site = web.TCPSite(self._runner, HOST, self._port)
|
||||
await site.start() # OSError (EADDRINUSE) propagates to caller
|
||||
self._publish_token()
|
||||
emit({"event": "ready", "data": {"port": self._port}})
|
||||
|
||||
async def stop(self) -> None:
|
||||
if self._runner is not None:
|
||||
await self._runner.cleanup()
|
||||
self._runner = None
|
||||
|
||||
|
||||
async def serve(
|
||||
service: VpnService,
|
||||
port: int = DEFAULT_PORT,
|
||||
token: str = "",
|
||||
token_file: Path | None = None,
|
||||
) -> ControlServer:
|
||||
server = ControlServer(service, port, token=token, token_file=token_file)
|
||||
await server.start()
|
||||
return server
|
||||
@@ -0,0 +1,6 @@
|
||||
"""Names used to identify processes and files owned by Ruh VPN."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
PREFIX = "ruh-vpn"
|
||||
TAG = "RUH_VPN_TAG"
|
||||
@@ -0,0 +1,192 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import Annotated, Literal, Optional, Union
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
|
||||
class _Base(BaseModel):
|
||||
model_config = ConfigDict(extra="allow", populate_by_name=True)
|
||||
|
||||
id: str = Field(default_factory=lambda: uuid.uuid4().hex[:12])
|
||||
name: str
|
||||
protocol: str
|
||||
|
||||
|
||||
class SSHServer(_Base):
|
||||
protocol: Literal["ssh"] = "ssh"
|
||||
host: str
|
||||
port: int = 22
|
||||
user: str
|
||||
password: Optional[str] = None
|
||||
keyFile: Optional[str] = None
|
||||
localPort: int = 11080
|
||||
|
||||
|
||||
class VlessServer(_Base):
|
||||
protocol: Literal["vless"] = "vless"
|
||||
address: str
|
||||
port: int
|
||||
uuid: str
|
||||
transport: str = "tcp"
|
||||
tls: bool = False
|
||||
sni: Optional[str] = None
|
||||
security: Optional[Literal["tls", "reality", "none"]] = None
|
||||
flow: Optional[str] = None
|
||||
fp: Optional[str] = None
|
||||
pbk: Optional[str] = None
|
||||
sid: Optional[str] = None
|
||||
# for ws/grpc/http transports
|
||||
path: Optional[str] = None
|
||||
host: Optional[str] = None
|
||||
serviceName: Optional[str] = None
|
||||
|
||||
|
||||
class VmessServer(_Base):
|
||||
protocol: Literal["vmess"] = "vmess"
|
||||
address: str
|
||||
port: int
|
||||
uuid: str
|
||||
alterId: int = 0
|
||||
security: str = "auto"
|
||||
transport: str = "tcp"
|
||||
tls: bool = False
|
||||
sni: Optional[str] = None
|
||||
path: Optional[str] = None
|
||||
host: Optional[str] = None
|
||||
|
||||
|
||||
class ShadowsocksServer(_Base):
|
||||
protocol: Literal["shadowsocks"] = "shadowsocks"
|
||||
address: str
|
||||
port: int
|
||||
method: str
|
||||
password: str
|
||||
|
||||
|
||||
class Socks5Server(_Base):
|
||||
protocol: Literal["socks5"] = "socks5"
|
||||
host: str
|
||||
port: int
|
||||
username: Optional[str] = None
|
||||
password: Optional[str] = None
|
||||
|
||||
|
||||
Server = Annotated[
|
||||
Union[SSHServer, VlessServer, VmessServer, ShadowsocksServer, Socks5Server],
|
||||
Field(discriminator="protocol"),
|
||||
]
|
||||
|
||||
|
||||
def parse_server(data: dict) -> Server:
|
||||
"""Parse a dict into one of the typed server models based on the protocol field."""
|
||||
protocol = (data.get("protocol") or "").lower()
|
||||
mapping = {
|
||||
"ssh": SSHServer,
|
||||
"vless": VlessServer,
|
||||
"vmess": VmessServer,
|
||||
"shadowsocks": ShadowsocksServer,
|
||||
"ss": ShadowsocksServer,
|
||||
"socks": Socks5Server,
|
||||
"socks5": Socks5Server,
|
||||
}
|
||||
cls = mapping.get(protocol)
|
||||
if cls is None:
|
||||
raise ValueError(f"Unsupported protocol: {protocol!r}")
|
||||
data = dict(data)
|
||||
if protocol == "ss":
|
||||
data["protocol"] = "shadowsocks"
|
||||
if protocol == "socks":
|
||||
data["protocol"] = "socks5"
|
||||
return cls.model_validate(data)
|
||||
|
||||
|
||||
def server_to_dict(server: BaseModel) -> dict:
|
||||
return server.model_dump(exclude_none=True)
|
||||
|
||||
|
||||
# Connection secrets never leave the backend: list RPCs strip them, and
|
||||
# UpdateServer treats an empty value as "keep the stored one".
|
||||
SENSITIVE_FIELDS = ("password", "uuid")
|
||||
|
||||
|
||||
def server_to_public_dict(server: BaseModel) -> dict:
|
||||
data = server.model_dump(exclude_none=True)
|
||||
for key in SENSITIVE_FIELDS:
|
||||
data.pop(key, None)
|
||||
return data
|
||||
|
||||
|
||||
class RoutingRule(BaseModel):
|
||||
"""User-defined routing rule.
|
||||
|
||||
type: force-proxy → match → proxy outbound
|
||||
direct → match → direct outbound
|
||||
block → match → block outbound
|
||||
pattern: one of
|
||||
- "example.com" exact domain
|
||||
- "*.example.com" domain suffix (matches example.com + subdomains)
|
||||
- "10.0.0.0/8" CIDR (v4 or v6)
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
id: str = Field(default_factory=lambda: uuid.uuid4().hex[:12])
|
||||
name: Optional[str] = None
|
||||
enabled: bool = True
|
||||
type: Literal["force-proxy", "direct", "block"] = "force-proxy"
|
||||
pattern: str
|
||||
|
||||
def to_singbox_rule(self) -> Optional[dict]:
|
||||
if not self.enabled or not self.pattern:
|
||||
return None
|
||||
rule: dict = {}
|
||||
if self.type == "block":
|
||||
rule["action"] = "reject"
|
||||
else:
|
||||
rule["outbound"] = "proxy" if self.type == "force-proxy" else "direct"
|
||||
pat = self.pattern.strip()
|
||||
if "/" in pat and not pat.startswith("*"):
|
||||
rule["ip_cidr"] = [pat]
|
||||
elif pat.startswith("*."):
|
||||
rule["domain_suffix"] = [pat[2:]]
|
||||
elif pat.startswith("*"):
|
||||
rule["domain_keyword"] = [pat.lstrip("*")]
|
||||
else:
|
||||
rule["domain"] = [pat]
|
||||
return rule
|
||||
|
||||
|
||||
class Settings(BaseModel):
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
activeServerId: Optional[str] = None
|
||||
mode: Literal["rules", "global"] = "rules"
|
||||
proxyMode: Literal["system", "tun"] = "system"
|
||||
autoStart: bool = False
|
||||
rulesPort: int = 11081
|
||||
globalPort: int = 11082
|
||||
transportPort: int = 11080
|
||||
refilterEnabled: bool = True
|
||||
healthCheckIntervalSec: int = 30
|
||||
killSwitchEnabled: bool = False
|
||||
clashApiPort: int = 11089
|
||||
showPingInBar: bool = True
|
||||
showTrafficInBar: bool = False
|
||||
activePresets: list[str] = Field(default_factory=lambda: ["ru"])
|
||||
|
||||
|
||||
class StatusInfo(BaseModel):
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
running: bool = False
|
||||
activeServerId: Optional[str] = None
|
||||
mode: str = "rules"
|
||||
proxyMode: str = "system"
|
||||
transportPort: int = 11080
|
||||
muxPort: Optional[int] = None
|
||||
pids: dict[str, int] = Field(default_factory=dict)
|
||||
message: Optional[str] = None
|
||||
status: str = "ok" # ok | degraded | failed | error
|
||||
reason: Optional[str] = None
|
||||
@@ -0,0 +1,331 @@
|
||||
"""Health monitoring: periodic TCP ping + traffic stats helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import socket
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import Awaitable, Callable, Optional
|
||||
|
||||
try:
|
||||
import aiohttp
|
||||
except ImportError: # pragma: no cover
|
||||
aiohttp = None # type: ignore[assignment]
|
||||
|
||||
try:
|
||||
from aiohttp_socks import ProxyConnector
|
||||
except ImportError: # pragma: no cover
|
||||
ProxyConnector = None # type: ignore[assignment]
|
||||
|
||||
SPEED_DOWN_URL = "https://speed.cloudflare.com/__down?bytes=10000000"
|
||||
SPEED_UP_URL = "https://speed.cloudflare.com/__up"
|
||||
SPEED_UP_BYTES = 4_000_000
|
||||
SPEED_TIMEOUT = 20.0
|
||||
|
||||
|
||||
def _connector_for(proxy_url: Optional[str]):
|
||||
"""Return an aiohttp connector. SOCKS5 proxy if given, else default."""
|
||||
if not proxy_url:
|
||||
return None
|
||||
if ProxyConnector is None:
|
||||
return None
|
||||
try:
|
||||
return ProxyConnector.from_url(proxy_url)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
async def tcp_ping(host: str, port: int, timeout: float = 5.0) -> Optional[int]:
|
||||
"""Open a TCP connection and return latency in ms, or None on failure."""
|
||||
loop = asyncio.get_running_loop()
|
||||
start = loop.time()
|
||||
try:
|
||||
fut = asyncio.open_connection(host, port)
|
||||
_, writer = await asyncio.wait_for(fut, timeout=timeout)
|
||||
latency_ms = int((loop.time() - start) * 1000)
|
||||
writer.close()
|
||||
try:
|
||||
await writer.wait_closed()
|
||||
except (ConnectionError, OSError):
|
||||
pass
|
||||
return latency_ms
|
||||
except (OSError, asyncio.TimeoutError):
|
||||
return None
|
||||
|
||||
|
||||
async def tcp_ping_samples(
|
||||
host: str,
|
||||
port: int,
|
||||
count: int = 4,
|
||||
timeout: float = 3.0,
|
||||
gap: float = 0.15,
|
||||
) -> list[int]:
|
||||
samples: list[int] = []
|
||||
for i in range(count):
|
||||
ms = await tcp_ping(host, port, timeout=timeout)
|
||||
if ms is not None:
|
||||
samples.append(ms)
|
||||
if i < count - 1:
|
||||
await asyncio.sleep(gap)
|
||||
return samples
|
||||
|
||||
|
||||
def compute_jitter(samples: list[int]) -> int:
|
||||
if len(samples) < 2:
|
||||
return 0
|
||||
diffs = [abs(samples[i] - samples[i - 1]) for i in range(1, len(samples))]
|
||||
return int(round(sum(diffs) / len(diffs)))
|
||||
|
||||
|
||||
async def measure_download_mbps(
|
||||
url: str = SPEED_DOWN_URL,
|
||||
timeout: float = SPEED_TIMEOUT,
|
||||
proxy_url: Optional[str] = None,
|
||||
) -> Optional[float]:
|
||||
"""Fetch URL and return throughput in Mbps (megabits/second).
|
||||
|
||||
If `proxy_url` is given (e.g. "socks5://127.0.0.1:11081"), traffic is
|
||||
routed through that SOCKS proxy so the measurement reflects the tunnel.
|
||||
"""
|
||||
if aiohttp is None:
|
||||
return None
|
||||
timeout_cfg = aiohttp.ClientTimeout(total=timeout, sock_connect=5.0)
|
||||
loop = asyncio.get_running_loop()
|
||||
connector = _connector_for(proxy_url)
|
||||
try:
|
||||
async with aiohttp.ClientSession(
|
||||
timeout=timeout_cfg, connector=connector
|
||||
) as session:
|
||||
async with session.get(url) as resp:
|
||||
if resp.status != 200:
|
||||
return None
|
||||
start = loop.time()
|
||||
total = 0
|
||||
async for chunk in resp.content.iter_chunked(65536):
|
||||
total += len(chunk)
|
||||
elapsed = max(loop.time() - start, 1e-6)
|
||||
if total <= 0:
|
||||
return None
|
||||
return round((total * 8.0) / elapsed / 1_000_000.0, 1)
|
||||
except (aiohttp.ClientError, asyncio.TimeoutError, OSError):
|
||||
return None
|
||||
|
||||
|
||||
async def measure_upload_mbps(
|
||||
url: str = SPEED_UP_URL,
|
||||
size_bytes: int = SPEED_UP_BYTES,
|
||||
timeout: float = SPEED_TIMEOUT,
|
||||
proxy_url: Optional[str] = None,
|
||||
) -> Optional[float]:
|
||||
if aiohttp is None:
|
||||
return None
|
||||
payload = b"\0" * size_bytes
|
||||
timeout_cfg = aiohttp.ClientTimeout(total=timeout, sock_connect=5.0)
|
||||
loop = asyncio.get_running_loop()
|
||||
connector = _connector_for(proxy_url)
|
||||
try:
|
||||
async with aiohttp.ClientSession(
|
||||
timeout=timeout_cfg, connector=connector
|
||||
) as session:
|
||||
start = loop.time()
|
||||
async with session.post(url, data=payload) as resp:
|
||||
# Read body to ensure full round-trip
|
||||
await resp.read()
|
||||
if resp.status >= 400:
|
||||
return None
|
||||
elapsed = max(loop.time() - start, 1e-6)
|
||||
return round((size_bytes * 8.0) / elapsed / 1_000_000.0, 1)
|
||||
except (aiohttp.ClientError, asyncio.TimeoutError, OSError):
|
||||
return None
|
||||
|
||||
|
||||
async def resolve_host(host: str) -> Optional[str]:
|
||||
loop = asyncio.get_running_loop()
|
||||
try:
|
||||
info = await loop.getaddrinfo(host, None, type=socket.SOCK_STREAM)
|
||||
for _, _, _, _, sockaddr in info:
|
||||
return sockaddr[0]
|
||||
except (socket.gaierror, OSError):
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def read_resolv_conf_nameservers(path: str = "/etc/resolv.conf") -> list[str]:
|
||||
out: list[str] = []
|
||||
try:
|
||||
with open(path, "r") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line.startswith("nameserver"):
|
||||
parts = line.split()
|
||||
if len(parts) >= 2:
|
||||
out.append(parts[1])
|
||||
except OSError:
|
||||
pass
|
||||
return out
|
||||
|
||||
|
||||
def check_dns_leak(running: bool, proxy_mode: str, mode: str) -> dict:
|
||||
"""Best-effort DNS leak check.
|
||||
|
||||
'leaking' is True when the proxy is up but system DNS is going to a
|
||||
nameserver that won't be routed through the proxy.
|
||||
|
||||
Heuristic:
|
||||
- TUN mode → all UDP/53 hits sing-box → NOT leaking (regardless of resolv.conf).
|
||||
- System+global → only HTTP/SOCKS goes through proxy; DNS to /etc/resolv.conf
|
||||
servers goes direct over the system → leaking.
|
||||
- System+rules → same as global from a DNS-leak standpoint → leaking.
|
||||
- Proxy not running → not running, "leaking" reported as N/A (false).
|
||||
"""
|
||||
nameservers = read_resolv_conf_nameservers()
|
||||
if not running:
|
||||
return {"leaking": False, "dns_servers": nameservers, "reason": "proxy not running"}
|
||||
if proxy_mode == "tun":
|
||||
return {"leaking": False, "dns_servers": nameservers, "reason": "TUN intercepts all DNS"}
|
||||
leaking = any(not (ns.startswith("127.") or ns == "::1") for ns in nameservers)
|
||||
return {
|
||||
"leaking": bool(leaking),
|
||||
"dns_servers": nameservers,
|
||||
"reason": (
|
||||
"system DNS bypasses the SOCKS proxy in system-proxy mode"
|
||||
if leaking
|
||||
else "all configured nameservers are local"
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class HealthState:
|
||||
latency_ms: int = -1
|
||||
jitter_ms: int = -1
|
||||
down_mbps: float = -1.0
|
||||
up_mbps: float = -1.0
|
||||
speed_taken_at: float = 0.0 # epoch seconds; 0 = never
|
||||
last_check: float = 0.0 # epoch seconds; 0 = never
|
||||
consecutive_failures: int = 0
|
||||
status: str = "ok" # ok | degraded | failed
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
if self.last_check:
|
||||
last_iso = datetime.fromtimestamp(self.last_check, tz=timezone.utc).isoformat()
|
||||
else:
|
||||
last_iso = ""
|
||||
if self.speed_taken_at:
|
||||
speed_iso = datetime.fromtimestamp(
|
||||
self.speed_taken_at, tz=timezone.utc
|
||||
).isoformat()
|
||||
else:
|
||||
speed_iso = ""
|
||||
return {
|
||||
"latency_ms": int(self.latency_ms),
|
||||
"jitter_ms": int(self.jitter_ms),
|
||||
"down_mbps": float(self.down_mbps),
|
||||
"up_mbps": float(self.up_mbps),
|
||||
"speed_taken_at": speed_iso,
|
||||
"last_check": last_iso,
|
||||
"consecutive_failures": int(self.consecutive_failures),
|
||||
"status": self.status,
|
||||
}
|
||||
|
||||
|
||||
class HealthMonitor:
|
||||
"""Periodic TCP ping to the transport server.
|
||||
|
||||
- One ping every `interval` seconds (default 30).
|
||||
- 1 failure → degraded; 3 consecutive → failed + on_failed callback fires once.
|
||||
- First successful ping after failure resets status to ok.
|
||||
"""
|
||||
|
||||
FAIL_THRESHOLD = 3
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
host: str,
|
||||
port: int,
|
||||
interval: float = 30.0,
|
||||
timeout: float = 5.0,
|
||||
on_failed: Optional[Callable[[], Awaitable[None]]] = None,
|
||||
) -> None:
|
||||
self.host = host
|
||||
self.port = port
|
||||
self.interval = interval
|
||||
self.timeout = timeout
|
||||
self._on_failed = on_failed
|
||||
self._task: Optional[asyncio.Task] = None
|
||||
self.state = HealthState()
|
||||
self._failed_emitted = False
|
||||
self._speed_task: Optional[asyncio.Task] = None
|
||||
|
||||
def start(self) -> None:
|
||||
if self._task and not self._task.done():
|
||||
return
|
||||
self._failed_emitted = False
|
||||
self._task = asyncio.create_task(self._loop())
|
||||
|
||||
async def stop(self) -> None:
|
||||
if self._task is None:
|
||||
return
|
||||
self._task.cancel()
|
||||
try:
|
||||
await self._task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._task = None
|
||||
|
||||
async def check_now(self) -> int:
|
||||
samples = await tcp_ping_samples(
|
||||
self.host, self.port, count=4, timeout=self.timeout, gap=0.1
|
||||
)
|
||||
self.state.last_check = time.time()
|
||||
if not samples:
|
||||
self.state.latency_ms = -1
|
||||
self.state.jitter_ms = -1
|
||||
self.state.consecutive_failures += 1
|
||||
if self.state.consecutive_failures >= self.FAIL_THRESHOLD:
|
||||
self.state.status = "failed"
|
||||
if not self._failed_emitted and self._on_failed:
|
||||
self._failed_emitted = True
|
||||
try:
|
||||
await self._on_failed()
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
self.state.status = "degraded"
|
||||
return -1
|
||||
latency = int(round(sum(samples) / len(samples)))
|
||||
self.state.latency_ms = latency
|
||||
self.state.jitter_ms = compute_jitter(samples)
|
||||
self.state.consecutive_failures = 0
|
||||
self.state.status = "ok"
|
||||
self._failed_emitted = False
|
||||
return latency
|
||||
|
||||
async def run_speed_test(self, proxy_url: Optional[str] = None) -> dict:
|
||||
"""Measure download + upload throughput. Updates state in-place.
|
||||
|
||||
When `proxy_url` is provided, traffic is routed through that proxy.
|
||||
"""
|
||||
down = await measure_download_mbps(proxy_url=proxy_url)
|
||||
up = await measure_upload_mbps(proxy_url=proxy_url)
|
||||
self.state.down_mbps = down if down is not None else -1.0
|
||||
self.state.up_mbps = up if up is not None else -1.0
|
||||
self.state.speed_taken_at = time.time()
|
||||
return {
|
||||
"down_mbps": self.state.down_mbps,
|
||||
"up_mbps": self.state.up_mbps,
|
||||
"ping_ms": int(self.state.latency_ms),
|
||||
"jitter_ms": int(self.state.jitter_ms),
|
||||
}
|
||||
|
||||
async def _loop(self) -> None:
|
||||
try:
|
||||
# First check immediately so GetHealth has real data quickly.
|
||||
await self.check_now()
|
||||
while True:
|
||||
await asyncio.sleep(self.interval)
|
||||
await self.check_now()
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
@@ -0,0 +1,145 @@
|
||||
"""Tail sing-box log files and forward each new line to a callback.
|
||||
|
||||
Each source is polled at a small interval; new bytes are split into complete
|
||||
lines (partial trailing data is buffered). ANSI escape codes are stripped and
|
||||
the line's log level is extracted when present (INFO/WARN/WARNING/ERROR/FATAL/
|
||||
DEBUG/TRACE), defaulting to "info".
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import Awaitable, Callable, Optional
|
||||
|
||||
ANSI_RE = re.compile(rb"\x1b\[[0-9;?]*[A-Za-z]")
|
||||
LEVEL_RE = re.compile(
|
||||
r"\b(TRACE|DEBUG|INFO|WARN(?:ING)?|ERROR|FATAL|PANIC)\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
LineCallback = Callable[[str, str, str], Awaitable[None]]
|
||||
# (source_tag, level, message)
|
||||
|
||||
|
||||
class _Source:
|
||||
def __init__(self, tag: str, path: Path) -> None:
|
||||
self.tag = tag
|
||||
self.path = path
|
||||
self.fd: Optional[int] = None
|
||||
self.buf = b""
|
||||
|
||||
def open(self) -> None:
|
||||
if self.fd is not None:
|
||||
return
|
||||
try:
|
||||
fd = os.open(str(self.path), os.O_RDONLY | os.O_NONBLOCK)
|
||||
except FileNotFoundError:
|
||||
return
|
||||
# seek to end so we don't emit historical content
|
||||
try:
|
||||
os.lseek(fd, 0, os.SEEK_END)
|
||||
except OSError:
|
||||
pass
|
||||
self.fd = fd
|
||||
|
||||
def close(self) -> None:
|
||||
if self.fd is not None:
|
||||
try:
|
||||
os.close(self.fd)
|
||||
except OSError:
|
||||
pass
|
||||
self.fd = None
|
||||
self.buf = b""
|
||||
|
||||
def read_lines(self) -> list[bytes]:
|
||||
if self.fd is None:
|
||||
self.open()
|
||||
if self.fd is None:
|
||||
return []
|
||||
try:
|
||||
chunk = os.read(self.fd, 65536)
|
||||
except BlockingIOError:
|
||||
return []
|
||||
except OSError:
|
||||
return []
|
||||
if not chunk:
|
||||
return []
|
||||
self.buf += chunk
|
||||
out: list[bytes] = []
|
||||
while True:
|
||||
nl = self.buf.find(b"\n")
|
||||
if nl < 0:
|
||||
break
|
||||
out.append(self.buf[:nl])
|
||||
self.buf = self.buf[nl + 1:]
|
||||
return out
|
||||
|
||||
|
||||
def parse_level(line: str) -> str:
|
||||
m = LEVEL_RE.search(line)
|
||||
if not m:
|
||||
return "info"
|
||||
lvl = m.group(1).lower()
|
||||
if lvl == "warning":
|
||||
return "warn"
|
||||
return lvl
|
||||
|
||||
|
||||
class LogStreamer:
|
||||
def __init__(self, callback: LineCallback, poll_interval: float = 0.5) -> None:
|
||||
self._cb = callback
|
||||
self._interval = poll_interval
|
||||
self._sources: dict[str, _Source] = {}
|
||||
self._task: Optional[asyncio.Task] = None
|
||||
|
||||
def add_source(self, tag: str, path: Path | str) -> None:
|
||||
p = Path(path)
|
||||
if tag in self._sources:
|
||||
return
|
||||
self._sources[tag] = _Source(tag, p)
|
||||
|
||||
def remove_source(self, tag: str) -> None:
|
||||
src = self._sources.pop(tag, None)
|
||||
if src:
|
||||
src.close()
|
||||
|
||||
def clear(self) -> None:
|
||||
for src in list(self._sources.values()):
|
||||
src.close()
|
||||
self._sources.clear()
|
||||
|
||||
def start(self) -> None:
|
||||
if self._task and not self._task.done():
|
||||
return
|
||||
self._task = asyncio.create_task(self._loop())
|
||||
|
||||
async def stop(self) -> None:
|
||||
if self._task is None:
|
||||
return
|
||||
self._task.cancel()
|
||||
try:
|
||||
await self._task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._task = None
|
||||
self.clear()
|
||||
|
||||
async def _loop(self) -> None:
|
||||
try:
|
||||
while True:
|
||||
await asyncio.sleep(self._interval)
|
||||
for src in list(self._sources.values()):
|
||||
for raw in src.read_lines():
|
||||
clean = ANSI_RE.sub(b"", raw).decode("utf-8", errors="replace").rstrip()
|
||||
if not clean:
|
||||
continue
|
||||
lvl = parse_level(clean)
|
||||
try:
|
||||
await self._cb(src.tag, lvl, clean)
|
||||
except Exception:
|
||||
pass
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
@@ -0,0 +1,95 @@
|
||||
"""Poll sing-box clash_api /connections endpoint for traffic stats."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Awaitable, Callable, Optional
|
||||
|
||||
import aiohttp
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrafficStats:
|
||||
bytes_sent: int = 0
|
||||
bytes_received: int = 0
|
||||
connection_count: int = 0
|
||||
started_at: float = field(default_factory=time.time)
|
||||
|
||||
def uptime(self) -> int:
|
||||
return max(0, int(time.time() - self.started_at))
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"bytes_sent": int(self.bytes_sent),
|
||||
"bytes_received": int(self.bytes_received),
|
||||
"uptime_seconds": self.uptime(),
|
||||
"connection_count": int(self.connection_count),
|
||||
}
|
||||
|
||||
|
||||
class TrafficMonitor:
|
||||
"""Polls /connections every `interval` seconds and tracks running totals.
|
||||
|
||||
Notes on the underlying API:
|
||||
sing-box clash_api /connections returns:
|
||||
{"downloadTotal": int, "uploadTotal": int, "connections": [...]}
|
||||
downloadTotal and uploadTotal are *since-mux-started* counters, so we
|
||||
can return them directly as bytes_received / bytes_sent.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_url: str,
|
||||
interval: float = 5.0,
|
||||
on_update: Optional[Callable[[dict], Awaitable[None]]] = None,
|
||||
) -> None:
|
||||
self.api_url = api_url.rstrip("/")
|
||||
self.interval = interval
|
||||
self._on_update = on_update
|
||||
self.stats = TrafficStats()
|
||||
self._task: Optional[asyncio.Task] = None
|
||||
|
||||
def start(self) -> None:
|
||||
self.stats = TrafficStats()
|
||||
if self._task and not self._task.done():
|
||||
return
|
||||
self._task = asyncio.create_task(self._loop())
|
||||
|
||||
async def stop(self) -> None:
|
||||
if self._task is None:
|
||||
return
|
||||
self._task.cancel()
|
||||
try:
|
||||
await self._task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._task = None
|
||||
|
||||
async def _poll_once(self) -> None:
|
||||
timeout = aiohttp.ClientTimeout(total=2.0)
|
||||
try:
|
||||
async with aiohttp.ClientSession(timeout=timeout) as session:
|
||||
async with session.get(f"{self.api_url}/connections") as resp:
|
||||
if resp.status != 200:
|
||||
return
|
||||
data = await resp.json(content_type=None)
|
||||
except (aiohttp.ClientError, asyncio.TimeoutError, OSError):
|
||||
return
|
||||
self.stats.bytes_sent = int(data.get("uploadTotal") or 0)
|
||||
self.stats.bytes_received = int(data.get("downloadTotal") or 0)
|
||||
self.stats.connection_count = len(data.get("connections") or [])
|
||||
|
||||
async def _loop(self) -> None:
|
||||
try:
|
||||
while True:
|
||||
await self._poll_once()
|
||||
if self._on_update:
|
||||
try:
|
||||
await self._on_update(self.stats.to_dict())
|
||||
except Exception:
|
||||
pass
|
||||
await asyncio.sleep(self.interval)
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
@@ -0,0 +1,25 @@
|
||||
"""Filesystem locations supplied by the Noctalia service entry."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _env_path(name: str, fallback: str) -> Path:
|
||||
return Path(os.path.expanduser(os.environ.get(name, fallback)))
|
||||
|
||||
|
||||
DATA_DIR = _env_path("RUH_VPN_DATA_DIR", "~/.local/share/ruh-vpn")
|
||||
RUNTIME_DIR = _env_path("RUH_VPN_RUNTIME_DIR", str(DATA_DIR / "runtime"))
|
||||
SINGBOX_DIR = DATA_DIR / "sing-box"
|
||||
|
||||
|
||||
def ensure_private_dir(path: Path) -> None:
|
||||
path.mkdir(mode=0o700, parents=True, exist_ok=True)
|
||||
path.chmod(0o700)
|
||||
|
||||
|
||||
def protect_file(path: Path) -> None:
|
||||
path.chmod(0o600)
|
||||
|
||||
@@ -0,0 +1,227 @@
|
||||
"""Helpers for translating user routing rules into sing-box rule entries."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from backend.models.server import RoutingRule
|
||||
|
||||
|
||||
# Country / region routing presets. Each one adds a pair of rule_set entries
|
||||
# (domains + IPs) and a single route rule that sends matches to the proxy.
|
||||
# Tags must be unique across active presets so sing-box doesn't reject the
|
||||
# config — the keys below were picked to avoid collisions.
|
||||
PRESETS: dict[str, dict[str, Any]] = {
|
||||
"ru": {
|
||||
"key": "ru",
|
||||
"name": "Russia",
|
||||
"flag": "🇷🇺",
|
||||
"description": "Re-filter list — sites and IPs blocked in Russia",
|
||||
"rule_sets": [
|
||||
{
|
||||
"tag": "refilter_domains",
|
||||
"type": "remote",
|
||||
"format": "binary",
|
||||
"url": "https://github.com/1andrevich/Re-filter-lists/releases/latest/download/ruleset-domain-refilter_domains.srs",
|
||||
"download_detour": "direct",
|
||||
},
|
||||
{
|
||||
"tag": "refilter_ipsum",
|
||||
"type": "remote",
|
||||
"format": "binary",
|
||||
"url": "https://github.com/1andrevich/Re-filter-lists/releases/latest/download/ruleset-ip-refilter_ipsum.srs",
|
||||
"download_detour": "direct",
|
||||
},
|
||||
],
|
||||
},
|
||||
# There is no "blocked in China" list: the GFW blocks foreign services, so
|
||||
# the standard bypass is the inverse — proxy everything geolocated OUTSIDE
|
||||
# China and let domestic traffic go direct. sing-geosite publishes its .srs
|
||||
# files on the `rule-set` branch, not as release assets.
|
||||
"cn": {
|
||||
"key": "cn",
|
||||
"name": "China",
|
||||
"flag": "🇨🇳",
|
||||
"description": "GFW bypass — foreign (non-Chinese) sites via VPN",
|
||||
"rule_sets": [
|
||||
{
|
||||
"tag": "geosite_noncn",
|
||||
"type": "remote",
|
||||
"format": "binary",
|
||||
"url": "https://raw.githubusercontent.com/SagerNet/sing-geosite/rule-set/geosite-geolocation-!cn.srs",
|
||||
"download_detour": "direct",
|
||||
},
|
||||
],
|
||||
},
|
||||
# geosite-sanctioned covers sites unavailable from Iran (state blocks and
|
||||
# foreign sanctions). geosite-ir would be the opposite — Iranian domestic
|
||||
# sites, which need no proxy. Same .srs-on-branch layout as sing-geosite.
|
||||
"ir": {
|
||||
"key": "ir",
|
||||
"name": "Iran",
|
||||
"flag": "🇮🇷",
|
||||
"description": "Sites unavailable from Iran (blocks and sanctions) via VPN",
|
||||
"rule_sets": [
|
||||
{
|
||||
"tag": "geosite_sanctioned",
|
||||
"type": "remote",
|
||||
"format": "binary",
|
||||
"url": "https://raw.githubusercontent.com/Chocolate4U/Iran-sing-box-rules/rule-set/geosite-sanctioned.srs",
|
||||
"download_detour": "direct",
|
||||
},
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def preset_rule_sets(active: list[str]) -> list[dict]:
|
||||
"""Return rule_set entries for the given active preset keys, deduped by tag."""
|
||||
seen: set[str] = set()
|
||||
out: list[dict] = []
|
||||
for key in active:
|
||||
preset = PRESETS.get(key)
|
||||
if not preset:
|
||||
continue
|
||||
for rs in preset["rule_sets"]:
|
||||
if rs["tag"] in seen:
|
||||
continue
|
||||
seen.add(rs["tag"])
|
||||
out.append(dict(rs))
|
||||
return out
|
||||
|
||||
|
||||
def preset_route_rules(active: list[str]) -> list[dict]:
|
||||
"""One route.rules entry per active preset routing its tags to 'proxy'."""
|
||||
out: list[dict] = []
|
||||
for key in active:
|
||||
preset = PRESETS.get(key)
|
||||
if not preset:
|
||||
continue
|
||||
tags = [rs["tag"] for rs in preset["rule_sets"]]
|
||||
if tags:
|
||||
out.append({"rule_set": tags, "outbound": "proxy"})
|
||||
return out
|
||||
|
||||
|
||||
def preset_domain_tags(active: list[str]) -> list[str]:
|
||||
"""Tags of rule_sets that match domains (used for proxy-DNS rule).
|
||||
|
||||
Heuristic: any tag containing 'domain' or 'site' is treated as domain-only.
|
||||
IPs don't help the DNS layer, so we skip them here.
|
||||
"""
|
||||
tags: list[str] = []
|
||||
for key in active:
|
||||
preset = PRESETS.get(key)
|
||||
if not preset:
|
||||
continue
|
||||
for rs in preset["rule_sets"]:
|
||||
t = rs["tag"]
|
||||
tl = t.lower()
|
||||
if "domain" in tl or "site" in tl:
|
||||
tags.append(t)
|
||||
return tags
|
||||
|
||||
# A hostname label: 1–63 chars, alphanumerics + hyphens (not at edges).
|
||||
# `*` is allowed as the leftmost label so wildcards like *.example.com work.
|
||||
_LABEL_RE = re.compile(r"^(?:\*|[A-Za-z0-9](?:[A-Za-z0-9-]{0,61}[A-Za-z0-9])?)$")
|
||||
|
||||
|
||||
def normalize_pattern(pattern: str) -> str:
|
||||
"""Clean a user-supplied pattern.
|
||||
|
||||
- Strips http:// and https:// prefixes (extracts hostname).
|
||||
- Strips any path / query / fragment from a URL-like input.
|
||||
- Strips trailing slashes and surrounding whitespace.
|
||||
Domain and CIDR forms pass through unchanged (case-folded for domains).
|
||||
"""
|
||||
p = (pattern or "").strip()
|
||||
if not p:
|
||||
return ""
|
||||
low = p.lower()
|
||||
if low.startswith("http://"):
|
||||
p = p[len("http://"):]
|
||||
elif low.startswith("https://"):
|
||||
p = p[len("https://"):]
|
||||
# Cut anything after the host: path, query, fragment.
|
||||
for sep in ("/", "?", "#"):
|
||||
# Don't cut the slash in CIDRs (digits on the right of '/').
|
||||
if sep == "/" and "/" in p:
|
||||
host, _, tail = p.partition("/")
|
||||
if tail and tail[0].isdigit() and host and (host[0].isdigit() or ":" in host):
|
||||
# Looks like a CIDR — keep as-is.
|
||||
continue
|
||||
p = host
|
||||
elif sep in p:
|
||||
p = p.split(sep, 1)[0]
|
||||
# Strip credentials and port (e.g. user:pass@host:443).
|
||||
if "@" in p:
|
||||
p = p.split("@", 1)[1]
|
||||
# Port: only strip when it's not part of an IPv6 literal.
|
||||
if p.count(":") == 1 and not p.startswith("["):
|
||||
p = p.split(":", 1)[0]
|
||||
return p.rstrip(".").lower()
|
||||
|
||||
|
||||
def validate_pattern(pattern: str) -> str:
|
||||
"""Validate and return the normalized pattern.
|
||||
|
||||
Raises ValueError when the pattern is not a recognized
|
||||
domain / wildcard domain / CIDR form.
|
||||
"""
|
||||
p = normalize_pattern(pattern)
|
||||
if not p:
|
||||
raise ValueError("pattern is empty")
|
||||
|
||||
# CIDR
|
||||
if "/" in p and not p.startswith("*"):
|
||||
try:
|
||||
ipaddress.ip_network(p, strict=False)
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"invalid CIDR: {p!r} ({exc})") from exc
|
||||
return p
|
||||
|
||||
# Bare IP without prefix is not allowed here (CIDR only)
|
||||
try:
|
||||
ipaddress.ip_address(p)
|
||||
raise ValueError(f"{p!r} is a bare IP; use CIDR (e.g. {p}/32)")
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
# Wildcard or domain
|
||||
labels = p.split(".")
|
||||
if not labels or any(not _LABEL_RE.match(label) for label in labels):
|
||||
raise ValueError(f"invalid domain pattern: {p!r}")
|
||||
# `*` may only appear as the leftmost label.
|
||||
if any(label == "*" for label in labels[1:]):
|
||||
raise ValueError(f"wildcard '*' only allowed as leftmost label: {p!r}")
|
||||
return p
|
||||
|
||||
|
||||
def classify_pattern(pattern: str) -> str:
|
||||
"""Return one of: 'cidr', 'wildcard', 'keyword', 'domain'."""
|
||||
p = pattern.strip()
|
||||
if not p:
|
||||
return "domain"
|
||||
if "/" in p and not p.startswith("*"):
|
||||
return "cidr"
|
||||
if p.startswith("*."):
|
||||
return "wildcard"
|
||||
if p.startswith("*") or p.endswith("*"):
|
||||
return "keyword"
|
||||
return "domain"
|
||||
|
||||
|
||||
def rules_to_singbox(rules: list[RoutingRule]) -> list[dict]:
|
||||
"""Convert a list of user rules into sing-box route.rules entries.
|
||||
|
||||
Skips disabled rules and rules with no pattern.
|
||||
Result preserves input order (first match wins in sing-box).
|
||||
"""
|
||||
out: list[dict] = []
|
||||
for r in rules:
|
||||
sr = r.to_singbox_rule()
|
||||
if sr is not None:
|
||||
out.append(sr)
|
||||
return out
|
||||
@@ -0,0 +1,151 @@
|
||||
"""nftables-backed kill switch.
|
||||
|
||||
Generates a self-contained inet table that drops all traffic except:
|
||||
- loopback
|
||||
- established / related connections
|
||||
- the noctalia-tun0 device (when present)
|
||||
- the proxy mux/transport ports on 127.0.0.1 (already covered by loopback)
|
||||
- explicit allowances for the active VPN server's host:port (so sing-box
|
||||
can dial out to it after the rules are installed)
|
||||
|
||||
The table is named "noctalia_killswitch" so removal/replacement is cheap and
|
||||
does not touch anyone else's nftables config.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import ipaddress
|
||||
import shutil
|
||||
import subprocess
|
||||
from typing import Optional
|
||||
|
||||
TABLE_NAME = "noctalia_killswitch"
|
||||
NFT_BIN = "/usr/sbin/nft"
|
||||
|
||||
|
||||
def _nft_path() -> str:
|
||||
return shutil.which("nft") or NFT_BIN
|
||||
|
||||
|
||||
def build_ruleset(
|
||||
server_ips: Optional[list[str]],
|
||||
server_port: Optional[int],
|
||||
tun_iface: str = "noctalia-tun0",
|
||||
extra_allow_tcp: Optional[list[int]] = None,
|
||||
) -> str:
|
||||
"""Build the nft ruleset text.
|
||||
|
||||
server_ips must be literal IP addresses (the caller resolves domain names
|
||||
beforehand). Every value is re-parsed through the ipaddress module and
|
||||
re-emitted in canonical form; anything that does not parse is dropped, so
|
||||
an untrusted server entry can never inject nft syntax into the ruleset,
|
||||
which runs with root privileges.
|
||||
"""
|
||||
port = int(server_port) if server_port else None
|
||||
tcp_ports = [int(p) for p in (extra_allow_tcp or [])]
|
||||
server_lines = ""
|
||||
for raw in server_ips or []:
|
||||
try:
|
||||
ip = ipaddress.ip_address(str(raw).strip())
|
||||
except ValueError:
|
||||
continue
|
||||
keyword = "ip6" if ip.version == 6 else "ip"
|
||||
if port:
|
||||
server_lines += f" {keyword} daddr {ip} tcp dport {port} accept\n"
|
||||
else:
|
||||
server_lines += f" {keyword} daddr {ip} accept\n"
|
||||
|
||||
tcp_port_line = ""
|
||||
if tcp_ports:
|
||||
ports = "{ " + ", ".join(str(p) for p in tcp_ports) + " }"
|
||||
tcp_port_line = f" tcp dport {ports} accept\n"
|
||||
|
||||
return (
|
||||
f"table inet {TABLE_NAME} {{\n"
|
||||
f" chain output {{\n"
|
||||
f" type filter hook output priority filter; policy drop;\n"
|
||||
f" oif \"lo\" accept\n"
|
||||
f" ct state established,related accept\n"
|
||||
f" oifname \"{tun_iface}\" accept\n"
|
||||
f" udp dport 53 accept\n"
|
||||
f" ip daddr 192.168.0.0/16 accept\n"
|
||||
f" ip daddr 10.0.0.0/8 accept\n"
|
||||
f" ip daddr 172.16.0.0/12 accept\n"
|
||||
f"{server_lines}{tcp_port_line}"
|
||||
f" }}\n"
|
||||
f" chain input {{\n"
|
||||
f" type filter hook input priority filter; policy drop;\n"
|
||||
f" iif \"lo\" accept\n"
|
||||
f" ct state established,related accept\n"
|
||||
f" iifname \"{tun_iface}\" accept\n"
|
||||
f" }}\n"
|
||||
f"}}\n"
|
||||
)
|
||||
|
||||
|
||||
async def _run_nft(args: list[str], input_text: Optional[str] = None) -> tuple[int, str]:
|
||||
nft = _nft_path()
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
nft,
|
||||
*args,
|
||||
stdin=subprocess.PIPE if input_text is not None else subprocess.DEVNULL,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
)
|
||||
out, _ = await proc.communicate(input_text.encode() if input_text else None)
|
||||
return proc.returncode or 0, out.decode("utf-8", errors="replace")
|
||||
|
||||
|
||||
async def _run_nft_via_pkexec(args: list[str], input_text: Optional[str] = None) -> tuple[int, str]:
|
||||
if shutil.which("pkexec") is None:
|
||||
return 1, "pkexec not available"
|
||||
nft = _nft_path()
|
||||
cmd = ["pkexec", nft, *args]
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
*cmd,
|
||||
stdin=subprocess.PIPE if input_text is not None else subprocess.DEVNULL,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
)
|
||||
out, _ = await proc.communicate(input_text.encode() if input_text else None)
|
||||
return proc.returncode or 0, out.decode("utf-8", errors="replace")
|
||||
|
||||
|
||||
async def apply(ruleset: str) -> tuple[bool, str]:
|
||||
"""Install (or replace) the kill switch ruleset. Returns (ok, message)."""
|
||||
# remove any prior version atomically before re-adding (idempotent)
|
||||
purge_cmd = f"delete table inet {TABLE_NAME}\n" + ruleset
|
||||
rc, out = await _run_nft(["-f", "-"], input_text=purge_cmd)
|
||||
if rc == 0:
|
||||
return True, "applied"
|
||||
# try just the add (no prior table)
|
||||
rc2, out2 = await _run_nft(["-f", "-"], input_text=ruleset)
|
||||
if rc2 == 0:
|
||||
return True, "applied"
|
||||
# fall back to pkexec
|
||||
rc3, out3 = await _run_nft_via_pkexec(["-f", "-"], input_text=purge_cmd)
|
||||
if rc3 == 0:
|
||||
return True, "applied via pkexec"
|
||||
rc4, out4 = await _run_nft_via_pkexec(["-f", "-"], input_text=ruleset)
|
||||
if rc4 == 0:
|
||||
return True, "applied via pkexec"
|
||||
return False, f"nft failed: {out2 or out}; pkexec: {out4 or out3}"
|
||||
|
||||
|
||||
async def remove() -> tuple[bool, str]:
|
||||
rc, out = await _run_nft(["delete", "table", "inet", TABLE_NAME])
|
||||
if rc == 0:
|
||||
return True, "removed"
|
||||
rc2, out2 = await _run_nft_via_pkexec(["delete", "table", "inet", TABLE_NAME])
|
||||
if rc2 == 0:
|
||||
return True, "removed via pkexec"
|
||||
# if the table doesn't exist, treat as success
|
||||
if "No such file or directory" in (out + out2) or "does not exist" in (out + out2):
|
||||
return True, "no table to remove"
|
||||
return False, f"nft failed: {out}; pkexec: {out2}"
|
||||
|
||||
|
||||
async def is_active() -> bool:
|
||||
rc, out = await _run_nft(["list", "table", "inet", TABLE_NAME])
|
||||
return rc == 0
|
||||
@@ -0,0 +1,48 @@
|
||||
"""Private copy of the sing-box binary used only for TUN mode.
|
||||
|
||||
CAP_NET_ADMIN is granted to a plugin-private copy under DATA_DIR/bin (a 0700
|
||||
directory) instead of the shared system binary, so the privilege never
|
||||
extends to other users or to sing-box invocations outside this plugin.
|
||||
|
||||
When the system binary changes, the copy is rewritten from scratch; a fresh
|
||||
file starts with no capabilities, so a stale copy never keeps the grant
|
||||
across sing-box upgrades.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
from backend.paths import DATA_DIR, ensure_private_dir
|
||||
|
||||
BIN_DIR = DATA_DIR / "bin"
|
||||
TUN_BIN = BIN_DIR / "sing-box-tun"
|
||||
|
||||
|
||||
def source_binary(singbox_bin: str) -> str:
|
||||
# setcap/getcap act on the real file, not a symlink (NixOS wraps binaries
|
||||
# in store symlinks, and setcap on the link fails).
|
||||
return os.path.realpath(singbox_bin)
|
||||
|
||||
|
||||
def ensure_copy(singbox_bin: str) -> tuple[str, bool]:
|
||||
"""Make sure the private copy exists and matches the system binary.
|
||||
|
||||
Returns (path to the copy, True if the copy was (re)created). Callers must
|
||||
treat a recreated copy as having no capabilities.
|
||||
"""
|
||||
src = Path(source_binary(singbox_bin))
|
||||
ensure_private_dir(BIN_DIR)
|
||||
st_src = src.stat()
|
||||
if TUN_BIN.exists():
|
||||
st_dst = TUN_BIN.stat()
|
||||
if st_dst.st_size == st_src.st_size and st_dst.st_mtime == st_src.st_mtime:
|
||||
return str(TUN_BIN), False
|
||||
# copy2 preserves mtime, which the staleness check above relies on
|
||||
tmp = TUN_BIN.with_name(TUN_BIN.name + ".tmp")
|
||||
shutil.copy2(src, tmp)
|
||||
tmp.chmod(0o700)
|
||||
os.replace(tmp, TUN_BIN)
|
||||
return str(TUN_BIN), True
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,329 @@
|
||||
"""Build sing-box JSON configs for transport / rules-mux / global-mux / TUN.
|
||||
|
||||
Each helper returns a dict that can be JSON-dumped straight into the matching
|
||||
the plugin data directory as <PREFIX>-{transport,rules,global,tun}.json.
|
||||
|
||||
Architecture (as proven by reference noctalia-rules.json / noctalia-global.json):
|
||||
|
||||
Transport layer → port 11080 (talks to remote VPN server)
|
||||
Rules mux → port 11081 (refilter rules → proxy, rest → direct)
|
||||
Global mux → port 11082 (everything → proxy)
|
||||
TUN → tun device; outbound = socks5 → 11081 or 11082
|
||||
|
||||
The TUN config never talks to the remote server directly — it always hops
|
||||
through one of the mux ports so we never create a routing loop.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from backend.models.server import RoutingRule, Server, SSHServer
|
||||
from backend.routing.rules import (
|
||||
preset_domain_tags,
|
||||
preset_route_rules,
|
||||
preset_rule_sets,
|
||||
)
|
||||
from backend.identity import PREFIX
|
||||
from backend.paths import SINGBOX_DIR
|
||||
from backend.singbox.transport import build_outbound
|
||||
|
||||
CONFIG_DIR = SINGBOX_DIR
|
||||
RULESET_CACHE_DIR = CONFIG_DIR # sing-box stores ruleset cache here
|
||||
RULES_DB = CONFIG_DIR / f"{PREFIX}-rules.db"
|
||||
|
||||
DEFAULT_LOG = {"level": "info", "timestamp": True}
|
||||
|
||||
PROXY_DNS_ADDR = "8.8.8.8"
|
||||
DIRECT_DNS_ADDR = "223.5.5.5"
|
||||
TUN_DNS_SERVER_NAME = "dns.google"
|
||||
|
||||
|
||||
def _dns_rules_from_user(rules: list) -> list[dict]:
|
||||
"""Translate user routing rules into DNS rules with matching server tags.
|
||||
|
||||
For each enabled rule:
|
||||
- extract matcher (domain / domain_suffix / domain_keyword / ip_cidr)
|
||||
- force-proxy → server: proxy-dns
|
||||
- direct → server: direct-dns
|
||||
- block → action: reject (no DNS lookup at all)
|
||||
"""
|
||||
out: list[dict] = []
|
||||
for r in rules:
|
||||
sr = r.to_singbox_rule() if hasattr(r, "to_singbox_rule") else None
|
||||
if not sr:
|
||||
continue
|
||||
dns_rule: dict = {}
|
||||
for k in ("domain", "domain_suffix", "domain_keyword", "ip_cidr"):
|
||||
if k in sr:
|
||||
dns_rule[k] = sr[k]
|
||||
if not dns_rule:
|
||||
continue
|
||||
if sr.get("action") == "reject":
|
||||
dns_rule["action"] = "reject"
|
||||
elif sr.get("outbound") == "proxy":
|
||||
dns_rule["server"] = "proxy-dns"
|
||||
else:
|
||||
dns_rule["server"] = "direct-dns"
|
||||
out.append(dns_rule)
|
||||
return out
|
||||
|
||||
|
||||
def _build_dns_rules(
|
||||
custom_rules: list,
|
||||
active_presets: list[str],
|
||||
default_proxy: bool,
|
||||
) -> dict:
|
||||
"""Return the dns section for a mux config.
|
||||
|
||||
default_proxy=True → unmatched DNS goes through proxy (global mode).
|
||||
default_proxy=False → unmatched DNS goes direct (rules mode).
|
||||
|
||||
For each active preset, domain-style rule_sets are routed via proxy-dns
|
||||
so DNS resolution for blocked sites doesn't leak to the direct resolver.
|
||||
"""
|
||||
servers = [
|
||||
{
|
||||
"type": "udp",
|
||||
"tag": "proxy-dns",
|
||||
"server": PROXY_DNS_ADDR,
|
||||
"server_port": 53,
|
||||
"detour": "proxy",
|
||||
},
|
||||
{
|
||||
"type": "udp",
|
||||
"tag": "direct-dns",
|
||||
"server": DIRECT_DNS_ADDR,
|
||||
"server_port": 53,
|
||||
},
|
||||
]
|
||||
rules = _dns_rules_from_user(custom_rules)
|
||||
if not default_proxy:
|
||||
dom_tags = preset_domain_tags(active_presets or [])
|
||||
if dom_tags:
|
||||
rules.append({"rule_set": dom_tags, "server": "proxy-dns"})
|
||||
return {
|
||||
"servers": servers,
|
||||
"rules": rules,
|
||||
"final": "proxy-dns" if default_proxy else "direct-dns",
|
||||
"strategy": "ipv4_only",
|
||||
}
|
||||
|
||||
|
||||
def build_transport_config(server: Server, listen_port: int = 11080) -> dict[str, Any]:
|
||||
"""Build sing-box config for the transport layer.
|
||||
|
||||
Listens on 127.0.0.1:listen_port (SOCKS5) and forwards through the
|
||||
server-specific outbound.
|
||||
|
||||
For SSH, this returns None — SSH is handled outside sing-box.
|
||||
"""
|
||||
if isinstance(server, SSHServer):
|
||||
raise ValueError(
|
||||
"SSH is handled directly by OpenSSH; do not build a sing-box transport config"
|
||||
)
|
||||
outbound = build_outbound(server, tag="proxy")
|
||||
return {
|
||||
"log": DEFAULT_LOG,
|
||||
"inbounds": [
|
||||
{
|
||||
"type": "socks",
|
||||
"tag": "in",
|
||||
"listen": "127.0.0.1",
|
||||
"listen_port": listen_port,
|
||||
"users": [],
|
||||
}
|
||||
],
|
||||
"outbounds": [
|
||||
outbound,
|
||||
{"type": "direct", "tag": "direct"},
|
||||
],
|
||||
"route": {"final": "proxy", "auto_detect_interface": True},
|
||||
}
|
||||
|
||||
|
||||
def build_rules_config(
|
||||
transport_port: int = 11080,
|
||||
listen_port: int = 11081,
|
||||
custom_rules: list[RoutingRule] | None = None,
|
||||
active_presets: list[str] | None = None,
|
||||
clash_api_port: int = 11089,
|
||||
) -> dict[str, Any]:
|
||||
"""Build the rules-mux config.
|
||||
|
||||
Listens on 127.0.0.1:listen_port (mixed inbound — accepts both SOCKS5 and
|
||||
HTTP), routes traffic per rules to either the upstream proxy (the transport
|
||||
listening on `transport_port`) or direct.
|
||||
|
||||
`active_presets` is a list of preset keys (e.g. ["ru"]). Each preset
|
||||
contributes its rule_set definitions and one route.rules entry that sends
|
||||
matches to the 'proxy' outbound. User custom_rules are placed first so they
|
||||
take precedence over preset rules (sing-box matches top-to-bottom).
|
||||
"""
|
||||
rules: list[dict[str, Any]] = []
|
||||
custom_rules = custom_rules or []
|
||||
active_presets = list(active_presets or [])
|
||||
|
||||
for r in custom_rules:
|
||||
if not r.enabled:
|
||||
continue
|
||||
sr = r.to_singbox_rule()
|
||||
if sr is not None:
|
||||
rules.append(sr)
|
||||
|
||||
rules.extend(preset_route_rules(active_presets))
|
||||
|
||||
rule_set = preset_rule_sets(active_presets)
|
||||
|
||||
route: dict[str, Any] = {
|
||||
"final": "direct",
|
||||
"auto_detect_interface": True,
|
||||
"default_domain_resolver": "direct-dns",
|
||||
"rules": rules,
|
||||
}
|
||||
if rule_set:
|
||||
route["rule_set"] = rule_set
|
||||
|
||||
return {
|
||||
"log": DEFAULT_LOG,
|
||||
"dns": _build_dns_rules(custom_rules, active_presets, default_proxy=False),
|
||||
"experimental": {
|
||||
"cache_file": {"enabled": True, "path": str(RULES_DB)},
|
||||
"clash_api": {"external_controller": f"127.0.0.1:{clash_api_port}"},
|
||||
},
|
||||
"inbounds": [
|
||||
{
|
||||
"type": "mixed",
|
||||
"tag": "in",
|
||||
"listen": "127.0.0.1",
|
||||
"listen_port": listen_port,
|
||||
}
|
||||
],
|
||||
"outbounds": [
|
||||
{"type": "direct", "tag": "direct"},
|
||||
{
|
||||
"type": "socks",
|
||||
"tag": "proxy",
|
||||
"server": "127.0.0.1",
|
||||
"server_port": transport_port,
|
||||
"version": "5",
|
||||
},
|
||||
],
|
||||
"route": route,
|
||||
}
|
||||
|
||||
|
||||
def build_global_config(
|
||||
transport_port: int = 11080,
|
||||
listen_port: int = 11082,
|
||||
clash_api_port: int | None = 11089,
|
||||
) -> dict[str, Any]:
|
||||
"""Build the global-mux config: everything → proxy."""
|
||||
experimental: dict[str, Any] = {}
|
||||
if clash_api_port is not None:
|
||||
experimental["clash_api"] = {"external_controller": f"127.0.0.1:{clash_api_port}"}
|
||||
return {
|
||||
"log": DEFAULT_LOG,
|
||||
"dns": _build_dns_rules([], active_presets=[], default_proxy=True),
|
||||
**({"experimental": experimental} if experimental else {}),
|
||||
"inbounds": [
|
||||
{
|
||||
"type": "mixed",
|
||||
"tag": "in",
|
||||
"listen": "127.0.0.1",
|
||||
"listen_port": listen_port,
|
||||
}
|
||||
],
|
||||
"outbounds": [
|
||||
{
|
||||
"type": "socks",
|
||||
"tag": "proxy",
|
||||
"server": "127.0.0.1",
|
||||
"server_port": transport_port,
|
||||
"version": "5",
|
||||
},
|
||||
{"type": "direct", "tag": "direct"},
|
||||
],
|
||||
"route": {
|
||||
"final": "proxy",
|
||||
"auto_detect_interface": True,
|
||||
"default_domain_resolver": "proxy-dns",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def build_tun_config(
|
||||
upstream_socks_port: int,
|
||||
interface_name: str = "noctalia-tun0",
|
||||
inet4_address: str = "172.19.0.1/30",
|
||||
route_exclude_addresses: list[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Build the TUN config.
|
||||
|
||||
The TUN outbound is a SOCKS5 client to 127.0.0.1:upstream_socks_port
|
||||
(either the rules mux on 11081 or the global mux on 11082). Private/LAN
|
||||
traffic goes direct so we don't black-hole local services.
|
||||
|
||||
sing-box exposes the second address of the TUN subnet (172.19.0.2 by
|
||||
default) to systemd-resolved. DNS must therefore be hijacked before the
|
||||
private-address rule, otherwise queries are sent direct to that synthetic
|
||||
address and immediately re-enter the TUN in a tight loop. DoH is used so
|
||||
SSH SOCKS transports, which cannot relay UDP, work as well.
|
||||
|
||||
``route_exclude_addresses`` contains the resolved transport endpoint(s).
|
||||
They must stay on the physical interface or an SSH/VPN transport would be
|
||||
captured by the TUN and recursively sent through itself.
|
||||
"""
|
||||
tun_inbound: dict[str, Any] = {
|
||||
"type": "tun",
|
||||
"tag": "tun-in",
|
||||
"interface_name": interface_name,
|
||||
"address": [inet4_address],
|
||||
"auto_route": True,
|
||||
"strict_route": True,
|
||||
"stack": "system",
|
||||
}
|
||||
if route_exclude_addresses:
|
||||
tun_inbound["route_exclude_address"] = route_exclude_addresses
|
||||
|
||||
return {
|
||||
"log": DEFAULT_LOG,
|
||||
"dns": {
|
||||
"servers": [
|
||||
{
|
||||
"type": "https",
|
||||
"tag": "tun-dns",
|
||||
"server": PROXY_DNS_ADDR,
|
||||
"server_port": 443,
|
||||
"path": "/dns-query",
|
||||
"tls": {
|
||||
"enabled": True,
|
||||
"server_name": TUN_DNS_SERVER_NAME,
|
||||
},
|
||||
"detour": "proxy",
|
||||
}
|
||||
],
|
||||
"final": "tun-dns",
|
||||
"strategy": "ipv4_only",
|
||||
},
|
||||
"inbounds": [tun_inbound],
|
||||
"outbounds": [
|
||||
{
|
||||
"type": "socks",
|
||||
"tag": "proxy",
|
||||
"server": "127.0.0.1",
|
||||
"server_port": upstream_socks_port,
|
||||
"version": "5",
|
||||
},
|
||||
{"type": "direct", "tag": "direct"},
|
||||
],
|
||||
"route": {
|
||||
"rules": [
|
||||
{"action": "sniff"},
|
||||
{"protocol": "dns", "action": "hijack-dns"},
|
||||
{"ip_is_private": True, "outbound": "direct"},
|
||||
],
|
||||
"final": "proxy",
|
||||
"auto_detect_interface": True,
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,309 @@
|
||||
"""Async start/stop/monitor of sing-box and ssh transport processes.
|
||||
|
||||
All managed processes are tagged via either:
|
||||
- ssh: <TAG>=1 environment variable
|
||||
- sing-box: filename pattern <PREFIX>-*.json passed as -c argument
|
||||
|
||||
This is intentionally narrow so pkill_zombies can use very specific patterns
|
||||
and never affect unrelated proxy processes.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import signal
|
||||
import subprocess
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Awaitable, Callable, Optional
|
||||
|
||||
import aiofiles
|
||||
|
||||
from backend.identity import PREFIX, TAG
|
||||
from backend.paths import DATA_DIR, RUNTIME_DIR, SINGBOX_DIR, ensure_private_dir, protect_file
|
||||
|
||||
# PATH first so NixOS and other non-FHS layouts work; /usr/bin only as a
|
||||
# last-resort guess when the backend env has a stripped PATH.
|
||||
SINGBOX_BIN = shutil.which("sing-box") or "/usr/bin/sing-box"
|
||||
SSHPASS_BIN = shutil.which("sshpass") or "/usr/bin/sshpass"
|
||||
SSH_BIN = shutil.which("ssh") or "/usr/bin/ssh"
|
||||
|
||||
SINGBOX_CONFIG_DIR = SINGBOX_DIR
|
||||
LOG_DIR = RUNTIME_DIR
|
||||
STATE_FILE = LOG_DIR / f"{PREFIX}.state.json"
|
||||
|
||||
# Plugin-owned known-hosts: accept-new records a server's key on first connect
|
||||
# and every later connect verifies it, so a changed key fails loudly instead of
|
||||
# being silently ignored (the old UserKnownHostsFile=/dev/null behavior).
|
||||
KNOWN_HOSTS_FILE = DATA_DIR / "known_hosts"
|
||||
|
||||
CONFIG_NAMES = {
|
||||
"transport": f"{PREFIX}-transport.json",
|
||||
"rules": f"{PREFIX}-rules.json",
|
||||
"global": f"{PREFIX}-global.json",
|
||||
"tun": f"{PREFIX}-tun.json",
|
||||
}
|
||||
|
||||
LOG_NAMES = {
|
||||
"transport": f"{PREFIX}-transport.log",
|
||||
"rules": f"{PREFIX}-rules.log",
|
||||
"global": f"{PREFIX}-global.log",
|
||||
"tun": f"{PREFIX}-tun.log",
|
||||
"ssh": f"{PREFIX}-ssh.log",
|
||||
}
|
||||
|
||||
# Must only ever match processes started with this plugin's identity.
|
||||
PKILL_PATTERNS = [
|
||||
f"ssh.*{TAG}=1",
|
||||
f"sing-box.*{PREFIX}-",
|
||||
]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ManagedProc:
|
||||
name: str # one of: transport, rules, global, tun, ssh
|
||||
proc: asyncio.subprocess.Process
|
||||
cmd: list[str]
|
||||
log_path: Path
|
||||
started_at: float = field(default_factory=time.time)
|
||||
|
||||
@property
|
||||
def pid(self) -> int:
|
||||
return self.proc.pid
|
||||
|
||||
def is_running(self) -> bool:
|
||||
return self.proc.returncode is None
|
||||
|
||||
|
||||
class ProcessManager:
|
||||
def __init__(self, logger: Optional[Callable[[str, str], None]] = None) -> None:
|
||||
self._procs: dict[str, ManagedProc] = {}
|
||||
self._monitor_task: Optional[asyncio.Task] = None
|
||||
self._monitor_cb: Optional[Callable[[str], Awaitable[None]]] = None
|
||||
self._log = logger or (lambda level, msg: None)
|
||||
ensure_private_dir(SINGBOX_CONFIG_DIR)
|
||||
ensure_private_dir(LOG_DIR)
|
||||
|
||||
# ----------------------------------------------------------------- config IO
|
||||
|
||||
async def write_config(self, name: str, config: dict[str, Any]) -> Path:
|
||||
if name not in CONFIG_NAMES:
|
||||
raise ValueError(f"Unknown sing-box config name: {name}")
|
||||
path = SINGBOX_CONFIG_DIR / CONFIG_NAMES[name]
|
||||
async with aiofiles.open(path, "w") as f:
|
||||
await f.write(json.dumps(config, indent=2))
|
||||
protect_file(path)
|
||||
return path
|
||||
|
||||
def config_path(self, name: str) -> Path:
|
||||
return SINGBOX_CONFIG_DIR / CONFIG_NAMES[name]
|
||||
|
||||
# ----------------------------------------------------------------- launch
|
||||
|
||||
async def start_singbox(self, name: str, binary: Optional[str] = None) -> ManagedProc:
|
||||
if name not in CONFIG_NAMES:
|
||||
raise ValueError(f"Unknown sing-box config name: {name}")
|
||||
if name in self._procs and self._procs[name].is_running():
|
||||
raise RuntimeError(f"sing-box '{name}' already running")
|
||||
config_path = self.config_path(name)
|
||||
if not config_path.exists():
|
||||
raise FileNotFoundError(f"Missing config file: {config_path}")
|
||||
log_path = LOG_DIR / LOG_NAMES[name]
|
||||
log_fh = open(log_path, "ab") # binary, append; sing-box writes structured text
|
||||
protect_file(log_path)
|
||||
cmd = [binary or SINGBOX_BIN, "run", "-c", str(config_path), "-D", str(SINGBOX_CONFIG_DIR)]
|
||||
self._log("info", f"start sing-box ({name}): {' '.join(cmd)}")
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
*cmd,
|
||||
stdout=log_fh,
|
||||
stderr=log_fh,
|
||||
stdin=subprocess.DEVNULL,
|
||||
start_new_session=True,
|
||||
)
|
||||
log_fh.close()
|
||||
managed = ManagedProc(name=name, proc=proc, cmd=cmd, log_path=log_path)
|
||||
self._procs[name] = managed
|
||||
return managed
|
||||
|
||||
async def start_ssh(
|
||||
self,
|
||||
host: str,
|
||||
port: int,
|
||||
user: str,
|
||||
local_port: int,
|
||||
password: Optional[str] = None,
|
||||
key_file: Optional[str] = None,
|
||||
) -> ManagedProc:
|
||||
if "ssh" in self._procs and self._procs["ssh"].is_running():
|
||||
raise RuntimeError("ssh transport already running")
|
||||
|
||||
log_path = LOG_DIR / LOG_NAMES["ssh"]
|
||||
log_fh = open(log_path, "ab")
|
||||
protect_file(log_path)
|
||||
|
||||
KNOWN_HOSTS_FILE.touch(mode=0o600, exist_ok=True)
|
||||
protect_file(KNOWN_HOSTS_FILE)
|
||||
|
||||
env = dict(os.environ)
|
||||
env[TAG] = "1"
|
||||
|
||||
common_ssh_opts = [
|
||||
"-N",
|
||||
"-D",
|
||||
f"127.0.0.1:{local_port}",
|
||||
"-o",
|
||||
"ExitOnForwardFailure=yes",
|
||||
"-o",
|
||||
"ServerAliveInterval=30",
|
||||
"-o",
|
||||
"ServerAliveCountMax=3",
|
||||
"-o",
|
||||
"StrictHostKeyChecking=accept-new",
|
||||
"-o",
|
||||
f"UserKnownHostsFile={KNOWN_HOSTS_FILE}",
|
||||
"-o",
|
||||
f"SetEnv={TAG}=1",
|
||||
"-o",
|
||||
f"SendEnv={TAG}",
|
||||
"-p",
|
||||
str(port),
|
||||
]
|
||||
|
||||
if password:
|
||||
cmd = [SSHPASS_BIN, "-e", SSH_BIN, *common_ssh_opts, f"{user}@{host}"]
|
||||
env["SSHPASS"] = password
|
||||
elif key_file:
|
||||
cmd = [SSH_BIN, *common_ssh_opts, "-i", key_file, f"{user}@{host}"]
|
||||
else:
|
||||
cmd = [SSH_BIN, *common_ssh_opts, f"{user}@{host}"]
|
||||
|
||||
self._log("info", f"start ssh transport to {user}@{host}:{port} -D {local_port}")
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
*cmd,
|
||||
stdout=log_fh,
|
||||
stderr=log_fh,
|
||||
stdin=subprocess.DEVNULL,
|
||||
env=env,
|
||||
start_new_session=True,
|
||||
)
|
||||
log_fh.close()
|
||||
managed = ManagedProc(name="ssh", proc=proc, cmd=cmd, log_path=log_path)
|
||||
self._procs["ssh"] = managed
|
||||
return managed
|
||||
|
||||
# ----------------------------------------------------------------- stop / monitor
|
||||
|
||||
async def stop(self, name: str, timeout: float = 3.0) -> None:
|
||||
managed = self._procs.get(name)
|
||||
if managed is None:
|
||||
return
|
||||
if managed.is_running():
|
||||
try:
|
||||
managed.proc.terminate()
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
try:
|
||||
await asyncio.wait_for(managed.proc.wait(), timeout=timeout)
|
||||
except asyncio.TimeoutError:
|
||||
try:
|
||||
managed.proc.kill()
|
||||
await managed.proc.wait()
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
self._procs.pop(name, None)
|
||||
|
||||
async def stop_all(self) -> None:
|
||||
await asyncio.gather(*(self.stop(n) for n in list(self._procs.keys())))
|
||||
await self.pkill_zombies()
|
||||
|
||||
async def pkill_zombies(self) -> None:
|
||||
"""Kill any leftover processes matching our narrow patterns."""
|
||||
for pattern in PKILL_PATTERNS:
|
||||
try:
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
"pkill", "-f", pattern,
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
)
|
||||
await proc.wait()
|
||||
except FileNotFoundError:
|
||||
return
|
||||
|
||||
# ----------------------------------------------------------------- introspection
|
||||
|
||||
def running_pids(self) -> dict[str, int]:
|
||||
return {n: m.pid for n, m in self._procs.items() if m.is_running()}
|
||||
|
||||
def running_names(self) -> list[str]:
|
||||
return [n for n, m in self._procs.items() if m.is_running()]
|
||||
|
||||
def is_running(self, name: str) -> bool:
|
||||
m = self._procs.get(name)
|
||||
return bool(m and m.is_running())
|
||||
|
||||
async def read_log_tail(self, name: str, max_bytes: int = 8192) -> str:
|
||||
log_path = LOG_DIR / LOG_NAMES.get(name, "")
|
||||
if not log_path.exists():
|
||||
return ""
|
||||
size = log_path.stat().st_size
|
||||
offset = max(0, size - max_bytes)
|
||||
async with aiofiles.open(log_path, "rb") as f:
|
||||
await f.seek(offset)
|
||||
data = await f.read()
|
||||
try:
|
||||
return data.decode("utf-8", errors="replace")
|
||||
except UnicodeDecodeError:
|
||||
return data.decode("latin-1", errors="replace")
|
||||
|
||||
# ----------------------------------------------------------------- monitor loop
|
||||
|
||||
def start_monitor(self, on_unexpected_exit: Callable[[str], Awaitable[None]]) -> None:
|
||||
self._monitor_cb = on_unexpected_exit
|
||||
if self._monitor_task and not self._monitor_task.done():
|
||||
return
|
||||
self._monitor_task = asyncio.create_task(self._monitor_loop())
|
||||
|
||||
async def stop_monitor(self) -> None:
|
||||
if self._monitor_task and not self._monitor_task.done():
|
||||
self._monitor_task.cancel()
|
||||
try:
|
||||
await self._monitor_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._monitor_task = None
|
||||
|
||||
async def _monitor_loop(self) -> None:
|
||||
try:
|
||||
while True:
|
||||
await asyncio.sleep(1.0)
|
||||
for name, m in list(self._procs.items()):
|
||||
if not m.is_running():
|
||||
rc = m.proc.returncode
|
||||
self._log("error", f"managed process '{name}' exited rc={rc}")
|
||||
self._procs.pop(name, None)
|
||||
if self._monitor_cb:
|
||||
try:
|
||||
await self._monitor_cb(name)
|
||||
except Exception as exc:
|
||||
self._log("error", f"monitor callback failed: {exc}")
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
|
||||
# ----------------------------------------------------------------- state file
|
||||
|
||||
async def write_state(self, state: dict[str, Any]) -> None:
|
||||
tmp = STATE_FILE.with_suffix(".json.tmp")
|
||||
async with aiofiles.open(tmp, "w") as f:
|
||||
await f.write(json.dumps(state, indent=2))
|
||||
protect_file(tmp)
|
||||
os.replace(tmp, STATE_FILE)
|
||||
|
||||
async def clear_state(self) -> None:
|
||||
try:
|
||||
STATE_FILE.unlink()
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
@@ -0,0 +1,147 @@
|
||||
"""Protocol-specific outbound builders for sing-box.
|
||||
|
||||
Each builder returns the outbound dict that goes into the sing-box "outbounds"
|
||||
list when configuring the transport layer (the layer that actually talks to the
|
||||
remote VPN server).
|
||||
|
||||
SSH is handled outside sing-box (via OpenSSH itself opening a SOCKS5
|
||||
listener on the local transport port), so it does NOT appear here.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from backend.models.server import (
|
||||
Server,
|
||||
ShadowsocksServer,
|
||||
Socks5Server,
|
||||
SSHServer,
|
||||
VlessServer,
|
||||
VmessServer,
|
||||
)
|
||||
|
||||
|
||||
def build_outbound(server: Server, tag: str = "proxy") -> dict[str, Any]:
|
||||
"""Return a sing-box outbound dict for the given server.
|
||||
|
||||
Raises ValueError for SSH (not a sing-box outbound) and for unsupported
|
||||
protocols.
|
||||
"""
|
||||
if isinstance(server, SSHServer):
|
||||
raise ValueError("SSH transport is handled outside sing-box")
|
||||
if isinstance(server, VlessServer):
|
||||
return _build_vless(server, tag)
|
||||
if isinstance(server, VmessServer):
|
||||
return _build_vmess(server, tag)
|
||||
if isinstance(server, ShadowsocksServer):
|
||||
return _build_shadowsocks(server, tag)
|
||||
if isinstance(server, Socks5Server):
|
||||
return _build_socks5(server, tag)
|
||||
raise ValueError(f"Unsupported server type: {type(server).__name__}")
|
||||
|
||||
|
||||
def _build_tls(server: VlessServer | VmessServer) -> dict[str, Any] | None:
|
||||
if not getattr(server, "tls", False) and getattr(server, "security", None) not in (
|
||||
"tls",
|
||||
"reality",
|
||||
):
|
||||
return None
|
||||
tls: dict[str, Any] = {"enabled": True}
|
||||
if server.sni:
|
||||
tls["server_name"] = server.sni
|
||||
fp = getattr(server, "fp", None)
|
||||
if fp:
|
||||
tls["utls"] = {"enabled": True, "fingerprint": fp}
|
||||
if getattr(server, "security", None) == "reality":
|
||||
pbk = getattr(server, "pbk", None) or ""
|
||||
sid = getattr(server, "sid", None) or ""
|
||||
tls["reality"] = {"enabled": True, "public_key": pbk, "short_id": sid}
|
||||
return tls
|
||||
|
||||
|
||||
def _build_transport(server: VlessServer | VmessServer) -> dict[str, Any] | None:
|
||||
t = (getattr(server, "transport", "tcp") or "tcp").lower()
|
||||
if t in ("tcp", "raw", ""):
|
||||
return None
|
||||
if t == "ws":
|
||||
out: dict[str, Any] = {"type": "ws"}
|
||||
if getattr(server, "path", None):
|
||||
out["path"] = server.path
|
||||
if getattr(server, "host", None):
|
||||
out["headers"] = {"Host": server.host}
|
||||
return out
|
||||
if t == "grpc":
|
||||
return {"type": "grpc", "service_name": getattr(server, "serviceName", "") or ""}
|
||||
if t == "http":
|
||||
out = {"type": "http"}
|
||||
if getattr(server, "path", None):
|
||||
out["path"] = server.path
|
||||
if getattr(server, "host", None):
|
||||
out["host"] = [server.host]
|
||||
return out
|
||||
return None
|
||||
|
||||
|
||||
def _build_vless(s: VlessServer, tag: str) -> dict[str, Any]:
|
||||
out: dict[str, Any] = {
|
||||
"type": "vless",
|
||||
"tag": tag,
|
||||
"server": s.address,
|
||||
"server_port": s.port,
|
||||
"uuid": s.uuid,
|
||||
}
|
||||
if s.flow:
|
||||
out["flow"] = s.flow
|
||||
tls = _build_tls(s)
|
||||
if tls:
|
||||
out["tls"] = tls
|
||||
tp = _build_transport(s)
|
||||
if tp:
|
||||
out["transport"] = tp
|
||||
return out
|
||||
|
||||
|
||||
def _build_vmess(s: VmessServer, tag: str) -> dict[str, Any]:
|
||||
out: dict[str, Any] = {
|
||||
"type": "vmess",
|
||||
"tag": tag,
|
||||
"server": s.address,
|
||||
"server_port": s.port,
|
||||
"uuid": s.uuid,
|
||||
"alter_id": s.alterId,
|
||||
"security": s.security or "auto",
|
||||
}
|
||||
tls = _build_tls(s)
|
||||
if tls:
|
||||
out["tls"] = tls
|
||||
tp = _build_transport(s)
|
||||
if tp:
|
||||
out["transport"] = tp
|
||||
return out
|
||||
|
||||
|
||||
def _build_shadowsocks(s: ShadowsocksServer, tag: str) -> dict[str, Any]:
|
||||
return {
|
||||
"type": "shadowsocks",
|
||||
"tag": tag,
|
||||
"server": s.address,
|
||||
"server_port": s.port,
|
||||
"method": s.method,
|
||||
"password": s.password,
|
||||
}
|
||||
|
||||
|
||||
def _build_socks5(s: Socks5Server, tag: str) -> dict[str, Any]:
|
||||
out: dict[str, Any] = {
|
||||
"type": "socks",
|
||||
"tag": tag,
|
||||
"server": s.host,
|
||||
"server_port": s.port,
|
||||
"version": "5",
|
||||
}
|
||||
if s.username:
|
||||
out["username"] = s.username
|
||||
if s.password:
|
||||
out["password"] = s.password
|
||||
return out
|
||||
@@ -0,0 +1,74 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import aiofiles
|
||||
|
||||
from backend.models.server import RoutingRule, Server, parse_server, server_to_dict
|
||||
from backend.paths import DATA_DIR, ensure_private_dir, protect_file
|
||||
|
||||
SERVERS_FILE = DATA_DIR / "servers.json"
|
||||
RULES_FILE = DATA_DIR / "rules.json"
|
||||
|
||||
|
||||
def ensure_dirs() -> None:
|
||||
ensure_private_dir(DATA_DIR)
|
||||
|
||||
|
||||
async def load_servers() -> list[Server]:
|
||||
ensure_dirs()
|
||||
if not SERVERS_FILE.exists():
|
||||
return []
|
||||
try:
|
||||
async with aiofiles.open(SERVERS_FILE, "r") as f:
|
||||
raw = await f.read()
|
||||
data = json.loads(raw or "[]")
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return []
|
||||
servers: list[Server] = []
|
||||
for entry in data:
|
||||
try:
|
||||
servers.append(parse_server(entry))
|
||||
except (ValueError, KeyError):
|
||||
continue
|
||||
return servers
|
||||
|
||||
|
||||
async def save_servers(servers: list) -> None:
|
||||
ensure_dirs()
|
||||
data = [server_to_dict(s) for s in servers]
|
||||
tmp = SERVERS_FILE.with_suffix(".json.tmp")
|
||||
async with aiofiles.open(tmp, "w") as f:
|
||||
await f.write(json.dumps(data, indent=2))
|
||||
protect_file(tmp)
|
||||
os.replace(tmp, SERVERS_FILE)
|
||||
|
||||
|
||||
async def load_rules() -> list[RoutingRule]:
|
||||
ensure_dirs()
|
||||
if not RULES_FILE.exists():
|
||||
return []
|
||||
try:
|
||||
async with aiofiles.open(RULES_FILE, "r") as f:
|
||||
raw = await f.read()
|
||||
data = json.loads(raw or "[]")
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return []
|
||||
rules: list[RoutingRule] = []
|
||||
for entry in data:
|
||||
try:
|
||||
rules.append(RoutingRule.model_validate(entry))
|
||||
except ValueError:
|
||||
continue
|
||||
return rules
|
||||
|
||||
|
||||
async def save_rules(rules: list[RoutingRule]) -> None:
|
||||
ensure_dirs()
|
||||
data = [r.model_dump(exclude_none=True) for r in rules]
|
||||
tmp = RULES_FILE.with_suffix(".json.tmp")
|
||||
async with aiofiles.open(tmp, "w") as f:
|
||||
await f.write(json.dumps(data, indent=2))
|
||||
protect_file(tmp)
|
||||
os.replace(tmp, RULES_FILE)
|
||||
@@ -0,0 +1,37 @@
|
||||
"""Persistence for subscription metadata."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import aiofiles
|
||||
|
||||
from backend.paths import DATA_DIR, ensure_private_dir, protect_file
|
||||
|
||||
SUBS_FILE = DATA_DIR / "subscriptions.json"
|
||||
|
||||
|
||||
def ensure_dirs() -> None:
|
||||
ensure_private_dir(DATA_DIR)
|
||||
|
||||
|
||||
async def load_subscriptions() -> list[dict]:
|
||||
ensure_dirs()
|
||||
if not SUBS_FILE.exists():
|
||||
return []
|
||||
try:
|
||||
async with aiofiles.open(SUBS_FILE, "r") as f:
|
||||
raw = await f.read()
|
||||
return json.loads(raw or "[]")
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return []
|
||||
|
||||
|
||||
async def save_subscriptions(subs: list[dict]) -> None:
|
||||
ensure_dirs()
|
||||
tmp = SUBS_FILE.with_suffix(".json.tmp")
|
||||
async with aiofiles.open(tmp, "w") as f:
|
||||
await f.write(json.dumps(subs, indent=2))
|
||||
protect_file(tmp)
|
||||
os.replace(tmp, SUBS_FILE)
|
||||
@@ -0,0 +1,157 @@
|
||||
"""Fetch subscription URLs, parse, import into VpnService."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import aiohttp
|
||||
|
||||
from backend.models.server import parse_server, server_to_dict
|
||||
from backend.storage.subscriptions import load_subscriptions, save_subscriptions
|
||||
from backend.subscription.parsers import parse_share_link, parse_subscription_body
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from backend.service.vpn_service import VpnService
|
||||
|
||||
AUTO_UPDATE_INTERVAL_SEC = 24 * 3600
|
||||
FETCH_TIMEOUT_SEC = 30
|
||||
USER_AGENT = "ruh-vpn/0.1 (subscription-fetcher)"
|
||||
|
||||
|
||||
class SubscriptionManager:
|
||||
def __init__(self, service: "VpnService") -> None:
|
||||
self._svc = service
|
||||
self._task: Optional[asyncio.Task] = None
|
||||
self._subs: list[dict] = []
|
||||
|
||||
async def bootstrap(self) -> None:
|
||||
self._subs = await load_subscriptions()
|
||||
|
||||
def start_auto_update(self) -> None:
|
||||
if self._task and not self._task.done():
|
||||
return
|
||||
self._task = asyncio.create_task(self._auto_loop())
|
||||
|
||||
async def stop(self) -> None:
|
||||
if self._task is None:
|
||||
return
|
||||
self._task.cancel()
|
||||
try:
|
||||
await self._task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._task = None
|
||||
|
||||
async def list_subs(self) -> list[dict]:
|
||||
return [dict(s) for s in self._subs]
|
||||
|
||||
async def add(self, url: str, name: str = "") -> bool:
|
||||
url = url.strip()
|
||||
if not url:
|
||||
return False
|
||||
if any(s["url"] == url for s in self._subs):
|
||||
return False
|
||||
entry = {
|
||||
"url": url,
|
||||
"name": name or url,
|
||||
"last_updated": 0,
|
||||
"server_count": 0,
|
||||
}
|
||||
self._subs.append(entry)
|
||||
await save_subscriptions(self._subs)
|
||||
return True
|
||||
|
||||
async def remove(self, url: str) -> bool:
|
||||
before = len(self._subs)
|
||||
self._subs = [s for s in self._subs if s["url"] != url]
|
||||
if len(self._subs) == before:
|
||||
return False
|
||||
await save_subscriptions(self._subs)
|
||||
return True
|
||||
|
||||
async def update(self, url: str) -> int:
|
||||
"""Fetch a single subscription URL and import its servers. Returns count."""
|
||||
for s in self._subs:
|
||||
if s["url"] == url:
|
||||
return await self._fetch_and_import(s)
|
||||
return 0
|
||||
|
||||
async def update_all(self) -> int:
|
||||
total = 0
|
||||
for s in list(self._subs):
|
||||
total += await self._fetch_and_import(s)
|
||||
return total
|
||||
|
||||
async def _fetch_and_import(self, sub: dict) -> int:
|
||||
try:
|
||||
body = await self._fetch(sub["url"])
|
||||
except Exception as exc:
|
||||
self._svc._log("error", f"subscription fetch failed for {sub['url']}: {exc}")
|
||||
return 0
|
||||
links = parse_subscription_body(body)
|
||||
imported = 0
|
||||
existing_keys = {self._server_key(s) for s in self._svc.state.servers}
|
||||
for link in links:
|
||||
entry = parse_share_link(link)
|
||||
if not entry:
|
||||
continue
|
||||
try:
|
||||
server = parse_server(entry)
|
||||
except Exception:
|
||||
continue
|
||||
key = self._server_key(server)
|
||||
if key in existing_keys:
|
||||
# update existing entry's fields by replacing with new id
|
||||
existing = next(
|
||||
(s for s in self._svc.state.servers if self._server_key(s) == key), None
|
||||
)
|
||||
if existing:
|
||||
entry["id"] = existing.id
|
||||
await self._svc.update_server(entry)
|
||||
continue
|
||||
await self._svc.add_server(entry)
|
||||
existing_keys.add(key)
|
||||
imported += 1
|
||||
sub["last_updated"] = int(time.time())
|
||||
sub["server_count"] = len(links)
|
||||
await save_subscriptions(self._subs)
|
||||
return imported
|
||||
|
||||
@staticmethod
|
||||
def _server_key(server) -> tuple:
|
||||
if isinstance(server, dict):
|
||||
proto = server.get("protocol", "")
|
||||
addr = server.get("address") or server.get("host") or ""
|
||||
port = server.get("port")
|
||||
secret = server.get("uuid") or server.get("password") or ""
|
||||
return (proto, addr, port, secret)
|
||||
proto = getattr(server, "protocol", "")
|
||||
addr = getattr(server, "address", None) or getattr(server, "host", None) or ""
|
||||
port = getattr(server, "port", None)
|
||||
secret = (
|
||||
getattr(server, "uuid", None)
|
||||
or getattr(server, "password", None)
|
||||
or ""
|
||||
)
|
||||
return (proto, addr, port, secret)
|
||||
|
||||
async def _fetch(self, url: str) -> str:
|
||||
timeout = aiohttp.ClientTimeout(total=FETCH_TIMEOUT_SEC)
|
||||
headers = {"User-Agent": USER_AGENT}
|
||||
async with aiohttp.ClientSession(timeout=timeout, headers=headers) as session:
|
||||
async with session.get(url) as resp:
|
||||
resp.raise_for_status()
|
||||
return await resp.text(errors="replace")
|
||||
|
||||
async def _auto_loop(self) -> None:
|
||||
try:
|
||||
while True:
|
||||
await asyncio.sleep(AUTO_UPDATE_INTERVAL_SEC)
|
||||
try:
|
||||
await self.update_all()
|
||||
except Exception as exc:
|
||||
self._svc._log("error", f"auto-update failed: {exc}")
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
@@ -0,0 +1,311 @@
|
||||
"""Parse share links (vless / vmess / ss / socks5 / sn) into server dicts.
|
||||
|
||||
The output dict shape matches `backend.models.server.parse_server` so it can be
|
||||
fed straight into VpnService.add_server.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import binascii
|
||||
import json
|
||||
import re
|
||||
import struct
|
||||
import urllib.parse as urlparse
|
||||
import uuid
|
||||
import zlib
|
||||
|
||||
|
||||
def _b64_decode_padded(data: str) -> bytes:
|
||||
data = data.strip().replace("\n", "").replace("\r", "")
|
||||
pad = "=" * (-len(data) % 4)
|
||||
try:
|
||||
return base64.urlsafe_b64decode(data + pad)
|
||||
except (binascii.Error, ValueError):
|
||||
try:
|
||||
return base64.b64decode(data + pad)
|
||||
except (binascii.Error, ValueError):
|
||||
return b""
|
||||
|
||||
|
||||
def parse_subscription_body(body: str) -> list[str]:
|
||||
"""Return a list of share-link strings from a raw subscription body.
|
||||
|
||||
Body may be:
|
||||
- Base64 of newline-separated share links (most common).
|
||||
- Plain text with newline-separated share links.
|
||||
"""
|
||||
body = body.strip()
|
||||
if not body:
|
||||
return []
|
||||
if "://" not in body:
|
||||
decoded = _b64_decode_padded(body)
|
||||
try:
|
||||
body = decoded.decode("utf-8", errors="replace")
|
||||
except UnicodeDecodeError:
|
||||
return []
|
||||
out: list[str] = []
|
||||
for line in body.splitlines():
|
||||
line = line.strip()
|
||||
if "://" in line:
|
||||
out.append(line)
|
||||
return out
|
||||
|
||||
|
||||
def parse_share_link(link: str) -> dict | None:
|
||||
link = link.strip()
|
||||
if link.startswith("vless://"):
|
||||
return _parse_vless(link)
|
||||
if link.startswith("vmess://"):
|
||||
return _parse_vmess(link)
|
||||
if link.startswith("ss://"):
|
||||
return _parse_ss(link)
|
||||
if link.startswith("socks5://") or link.startswith("socks://"):
|
||||
return _parse_socks5(link)
|
||||
if link.startswith("sn://"):
|
||||
return _parse_sn(link)
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- sn:// links
|
||||
#
|
||||
# sn://<type>?<urlsafe-base64 of zlib(payload)>. The payload is a binary record,
|
||||
# not JSON: strings are stored raw with the high bit set on their LAST byte
|
||||
# (so "192.0.2." + 0xb1 reads as "192.0.2.1"), and numbers are 32-bit LE.
|
||||
#
|
||||
# This layout was derived by inspection of a working ssh link — no public spec
|
||||
# was found for it, and it is NOT nekoray's (its repositories contain no "sn://"
|
||||
# and it has no ssh profile type). The reading was confirmed field by field
|
||||
# against a real link: port came out as exactly 22 and the user as "root", and
|
||||
# the owner verified the decoded password character for character.
|
||||
#
|
||||
# Because the format is inferred rather than specified, everything here is
|
||||
# strict: only the "ssh" type is accepted, every field must survive validation,
|
||||
# and anything unexpected returns None so the caller reports an unsupported link
|
||||
# instead of silently creating a wrong server. Two trailing fields (an int and
|
||||
# what looks like a UTF-8 remark) do not fit the scheme and are ignored — the
|
||||
# name is taken from the host instead.
|
||||
|
||||
|
||||
def _sn_read_string(buf: bytes, i: int) -> tuple[str, int] | None:
|
||||
"""Read one high-bit-terminated string starting at `i`."""
|
||||
out = bytearray()
|
||||
while i < len(buf):
|
||||
c = buf[i]
|
||||
i += 1
|
||||
if c & 0x80:
|
||||
out.append(c & 0x7F)
|
||||
try:
|
||||
return out.decode("utf-8"), i
|
||||
except UnicodeDecodeError:
|
||||
return None
|
||||
out.append(c)
|
||||
return None # ran off the end without a terminator
|
||||
|
||||
|
||||
def _sn_read_u32(buf: bytes, i: int) -> tuple[int, int] | None:
|
||||
if i + 4 > len(buf):
|
||||
return None
|
||||
return struct.unpack_from("<I", buf, i)[0], i + 4
|
||||
|
||||
|
||||
_SN_PRINTABLE = re.compile(r"^[\x20-\x7e]+$")
|
||||
|
||||
|
||||
def _parse_sn(link: str) -> dict | None:
|
||||
try:
|
||||
kind, payload = link[len("sn://"):].split("?", 1)
|
||||
except ValueError:
|
||||
return None
|
||||
if kind != "ssh":
|
||||
return None # only type verified against a real link
|
||||
|
||||
raw = _b64_decode_padded(payload)
|
||||
if not raw:
|
||||
return None
|
||||
try:
|
||||
buf = zlib.decompress(raw)
|
||||
except zlib.error:
|
||||
return None
|
||||
|
||||
i = 4 # leading u32, always 0 in the sample; purpose unknown
|
||||
host_r = _sn_read_string(buf, i)
|
||||
if not host_r:
|
||||
return None
|
||||
host, i = host_r
|
||||
port_r = _sn_read_u32(buf, i)
|
||||
if not port_r:
|
||||
return None
|
||||
port, i = port_r
|
||||
user_r = _sn_read_string(buf, i)
|
||||
if not user_r:
|
||||
return None
|
||||
user, i = user_r
|
||||
# Unknown u32 between user and password (1 in the sample; possibly an auth
|
||||
# mode). Not trusted for anything.
|
||||
skip = _sn_read_u32(buf, i)
|
||||
if not skip:
|
||||
return None
|
||||
i = skip[1]
|
||||
pw_r = _sn_read_string(buf, i)
|
||||
if not pw_r:
|
||||
return None
|
||||
password, _ = pw_r
|
||||
|
||||
if not (0 < port < 65536):
|
||||
return None
|
||||
for value in (host, user, password):
|
||||
if not value or not _SN_PRINTABLE.match(value):
|
||||
return None
|
||||
|
||||
return {
|
||||
"protocol": "ssh",
|
||||
"name": host,
|
||||
"host": host,
|
||||
"port": port,
|
||||
"user": user,
|
||||
"password": password,
|
||||
}
|
||||
|
||||
|
||||
def _decode_name(fragment: str) -> str:
|
||||
return urlparse.unquote(fragment or "").strip() or "imported"
|
||||
|
||||
|
||||
def _parse_vless(link: str) -> dict | None:
|
||||
parsed = urlparse.urlparse(link)
|
||||
if not parsed.username or not parsed.hostname or not parsed.port:
|
||||
return None
|
||||
q = urlparse.parse_qs(parsed.query)
|
||||
|
||||
def _q(k: str, default: str = "") -> str:
|
||||
return (q.get(k, [default]) or [default])[0]
|
||||
|
||||
out: dict = {
|
||||
"name": _decode_name(parsed.fragment),
|
||||
"protocol": "vless",
|
||||
"address": parsed.hostname,
|
||||
"port": int(parsed.port),
|
||||
"uuid": parsed.username,
|
||||
"transport": _q("type", "tcp") or "tcp",
|
||||
}
|
||||
sec = _q("security", "")
|
||||
out["security"] = sec if sec in ("tls", "reality", "none") else None
|
||||
out["tls"] = bool(sec in ("tls", "reality"))
|
||||
sni = _q("sni") or _q("host")
|
||||
if sni:
|
||||
out["sni"] = sni
|
||||
for src, dst in [
|
||||
("flow", "flow"),
|
||||
("fp", "fp"),
|
||||
("pbk", "pbk"),
|
||||
("sid", "sid"),
|
||||
("path", "path"),
|
||||
("serviceName", "serviceName"),
|
||||
]:
|
||||
v = _q(src)
|
||||
if v:
|
||||
out[dst] = v
|
||||
out["id"] = _gen_id("vless", out["address"], out["port"], out["uuid"])
|
||||
return {k: v for k, v in out.items() if v is not None}
|
||||
|
||||
|
||||
def _parse_vmess(link: str) -> dict | None:
|
||||
payload = link[len("vmess://"):]
|
||||
decoded = _b64_decode_padded(payload)
|
||||
if not decoded:
|
||||
return None
|
||||
try:
|
||||
obj = json.loads(decoded.decode("utf-8", errors="replace"))
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return None
|
||||
addr = obj.get("add")
|
||||
port = obj.get("port")
|
||||
uuid_ = obj.get("id")
|
||||
if not addr or not port or not uuid_:
|
||||
return None
|
||||
out = {
|
||||
"name": obj.get("ps") or "imported",
|
||||
"protocol": "vmess",
|
||||
"address": addr,
|
||||
"port": int(port),
|
||||
"uuid": uuid_,
|
||||
"alterId": int(obj.get("aid") or 0),
|
||||
"security": obj.get("scy") or "auto",
|
||||
"transport": obj.get("net") or "tcp",
|
||||
"tls": (obj.get("tls") == "tls"),
|
||||
}
|
||||
if obj.get("sni") or obj.get("host"):
|
||||
out["sni"] = obj.get("sni") or obj.get("host")
|
||||
if obj.get("path"):
|
||||
out["path"] = obj["path"]
|
||||
if obj.get("host"):
|
||||
out["host"] = obj["host"]
|
||||
out["id"] = _gen_id("vmess", out["address"], out["port"], out["uuid"])
|
||||
return out
|
||||
|
||||
|
||||
def _parse_ss(link: str) -> dict | None:
|
||||
# Two common forms:
|
||||
# ss://base64(method:password)@host:port#name
|
||||
# ss://base64(method:password@host:port)#name
|
||||
rest = link[len("ss://"):]
|
||||
frag = ""
|
||||
if "#" in rest:
|
||||
rest, frag = rest.split("#", 1)
|
||||
name = _decode_name(frag)
|
||||
|
||||
method: str | None = None
|
||||
password: str | None = None
|
||||
host: str | None = None
|
||||
port: int | None = None
|
||||
|
||||
if "@" in rest:
|
||||
creds_b64, host_part = rest.rsplit("@", 1)
|
||||
creds = _b64_decode_padded(creds_b64).decode("utf-8", errors="replace")
|
||||
if ":" in creds:
|
||||
method, password = creds.split(":", 1)
|
||||
if ":" in host_part:
|
||||
h, p = host_part.rsplit(":", 1)
|
||||
host = h
|
||||
try:
|
||||
port = int(p)
|
||||
except ValueError:
|
||||
pass
|
||||
else:
|
||||
whole = _b64_decode_padded(rest).decode("utf-8", errors="replace")
|
||||
m = re.match(r"^([^:]+):([^@]+)@([^:]+):(\d+)$", whole)
|
||||
if m:
|
||||
method, password, host, port = m.group(1), m.group(2), m.group(3), int(m.group(4))
|
||||
|
||||
if not method or not password or not host or not port:
|
||||
return None
|
||||
return {
|
||||
"id": _gen_id("ss", host, port, password),
|
||||
"name": name,
|
||||
"protocol": "shadowsocks",
|
||||
"address": host,
|
||||
"port": port,
|
||||
"method": method,
|
||||
"password": password,
|
||||
}
|
||||
|
||||
|
||||
def _parse_socks5(link: str) -> dict | None:
|
||||
parsed = urlparse.urlparse(link)
|
||||
if not parsed.hostname or not parsed.port:
|
||||
return None
|
||||
return {
|
||||
"id": _gen_id("socks5", parsed.hostname, parsed.port, parsed.username or ""),
|
||||
"name": _decode_name(parsed.fragment),
|
||||
"protocol": "socks5",
|
||||
"host": parsed.hostname,
|
||||
"port": int(parsed.port),
|
||||
"username": parsed.username,
|
||||
"password": parsed.password,
|
||||
}
|
||||
|
||||
|
||||
def _gen_id(proto: str, host: str, port: int, secret: str) -> str:
|
||||
h = f"{proto}|{host}|{port}|{secret}".encode("utf-8")
|
||||
return uuid.uuid5(uuid.NAMESPACE_URL, h.decode("utf-8")).hex[:12]
|
||||
Reference in New Issue
Block a user