Files
singbox_glue/src/vpn_egressctl/policy.py
T

268 lines
10 KiB
Python

from __future__ import annotations
import ipaddress
import json
import re
from dataclasses import dataclass
from pathlib import Path, PurePosixPath
from typing import Any
from urllib.parse import urlsplit
from .errors import ValidationError
from .version import SUPPORTED_SING_BOX_VERSION
@dataclass(frozen=True, slots=True)
class SingBoxPolicy:
binary: str
config_path: str
service: str
required_version: str
@dataclass(frozen=True, slots=True)
class RuntimePolicy:
uri_path: str
state_dir: str
lock_path: str
backup_keep: int
@dataclass(frozen=True, slots=True)
class NetworkPolicy:
upstream_interface: str
vpn_lan_interface: str
tun_name: str
tun_address: str
mtu: int
route_exclude_address: tuple[str, ...]
iproute2_table_index: int
iproute2_rule_index: int
auto_redirect_input_mark: str
auto_redirect_output_mark: str
auto_redirect_reset_mark: str
auto_redirect_nfqueue: int
auto_redirect_fallback_rule_index: int
@dataclass(frozen=True, slots=True)
class DnsPolicy:
bootstrap_server: str
bootstrap_port: int
remote_server: str
remote_port: int
remote_path: str
remote_tls_server_name: str
strategy: str
@dataclass(frozen=True, slots=True)
class BandwidthPolicy:
up_mbps: int
down_mbps: int
@dataclass(frozen=True, slots=True)
class HealthcheckPolicy:
url: str | None
timeout_seconds: float
settle_seconds: float
expected_status: int
body_contains: str | None
@dataclass(frozen=True, slots=True)
class Policy:
schema_version: int
sing_box: SingBoxPolicy
runtime: RuntimePolicy
network: NetworkPolicy
dns: DnsPolicy
bandwidth: BandwidthPolicy
healthcheck: HealthcheckPolicy
_EXPECTED: dict[str, set[str]] = {
"root": {"schema_version", "sing_box", "runtime", "network", "dns", "bandwidth", "healthcheck"},
"sing_box": {"binary", "config_path", "service", "required_version"},
"runtime": {"uri_path", "state_dir", "lock_path", "backup_keep"},
"network": {
"upstream_interface", "vpn_lan_interface", "tun_name", "tun_address", "mtu",
"route_exclude_address", "iproute2_table_index", "iproute2_rule_index",
"auto_redirect_input_mark", "auto_redirect_output_mark", "auto_redirect_reset_mark",
"auto_redirect_nfqueue", "auto_redirect_fallback_rule_index",
},
"dns": {"bootstrap_server", "bootstrap_port", "remote_server", "remote_port", "remote_path", "remote_tls_server_name", "strategy"},
"bandwidth": {"up_mbps", "down_mbps"},
"healthcheck": {"url", "timeout_seconds", "settle_seconds", "expected_status", "body_contains"},
}
def _mapping(value: Any, label: str) -> dict[str, Any]:
if not isinstance(value, dict):
raise ValidationError(f"Policy section {label} must be an object")
unknown = set(value) - _EXPECTED[label]
missing = _EXPECTED[label] - set(value)
if unknown:
raise ValidationError(f"Unknown policy keys in {label}: {', '.join(sorted(unknown))}")
if missing:
raise ValidationError(f"Missing policy keys in {label}: {', '.join(sorted(missing))}")
return value
def _string(mapping: dict[str, Any], key: str, *, nonempty: bool = True) -> str:
value = mapping[key]
if not isinstance(value, str) or (nonempty and not value):
raise ValidationError(f"Policy value {key} must be a non-empty string")
return value
def _integer(mapping: dict[str, Any], key: str, minimum: int, maximum: int) -> int:
value = mapping[key]
if isinstance(value, bool) or not isinstance(value, int) or not minimum <= value <= maximum:
raise ValidationError(f"Policy value {key} must be in {minimum}..{maximum}")
return value
def _absolute(value: str, label: str) -> str:
path = PurePosixPath(value)
if not path.is_absolute() or ".." in path.parts:
raise ValidationError(f"Policy path {label} must be an absolute normalised POSIX path")
return str(path)
def _interface(value: str, label: str) -> str:
if not re.fullmatch(r"[A-Za-z0-9_.:-]{1,15}", value):
raise ValidationError(f"Invalid interface name in {label}")
return value
def _mark(value: str, label: str) -> str:
if not re.fullmatch(r"0x[0-9a-fA-F]{1,8}", value):
raise ValidationError(f"Invalid hexadecimal mark in {label}")
return "0x" + value[2:].lower()
def load_policy(path: str | Path) -> Policy:
try:
with Path(path).open("r", encoding="utf-8") as stream:
raw = json.load(stream)
except (OSError, json.JSONDecodeError) as exc:
raise ValidationError(f"Cannot read policy file: {path}") from exc
root = _mapping(raw, "root")
if root["schema_version"] != 1:
raise ValidationError("Only policy schema_version=1 is supported")
sb = _mapping(root["sing_box"], "sing_box")
runtime = _mapping(root["runtime"], "runtime")
network = _mapping(root["network"], "network")
dns = _mapping(root["dns"], "dns")
bandwidth = _mapping(root["bandwidth"], "bandwidth")
health = _mapping(root["healthcheck"], "healthcheck")
required_version = _string(sb, "required_version")
if required_version != SUPPORTED_SING_BOX_VERSION:
raise ValidationError(
f"This release requires sing-box version exactly {SUPPORTED_SING_BOX_VERSION}"
)
service = _string(sb, "service")
if not re.fullmatch(r"[A-Za-z0-9@_.:-]+\.service", service):
raise ValidationError("Invalid sing-box systemd service name")
exclusions_raw = network["route_exclude_address"]
if not isinstance(exclusions_raw, list) or not exclusions_raw:
raise ValidationError("route_exclude_address must be a non-empty array")
exclusions: list[str] = []
for value in exclusions_raw:
if not isinstance(value, str):
raise ValidationError("route_exclude_address entries must be strings")
try:
parsed = ipaddress.ip_network(value, strict=True)
except ValueError as exc:
raise ValidationError(f"Invalid route exclusion: {value}") from exc
exclusions.append(str(parsed))
tun_address = _string(network, "tun_address")
try:
ipaddress.ip_interface(tun_address)
except ValueError as exc:
raise ValidationError("Invalid tun_address") from exc
bootstrap = _string(dns, "bootstrap_server")
remote = _string(dns, "remote_server")
try:
ipaddress.ip_address(bootstrap)
ipaddress.ip_address(remote)
except ValueError as exc:
raise ValidationError("DNS bootstrap and remote servers must be IP addresses") from exc
if dns["strategy"] != "ipv4_only":
raise ValidationError("Only DNS strategy ipv4_only is supported")
url = health["url"]
body_contains = health["body_contains"]
if url is not None and (not isinstance(url, str) or not url.startswith("https://")):
raise ValidationError("healthcheck.url must be null or an https:// URL")
if url is not None:
parsed_url = urlsplit(url)
if not parsed_url.hostname or parsed_url.username is not None or parsed_url.password is not None or parsed_url.fragment:
raise ValidationError("healthcheck.url must not contain credentials or a fragment")
if not _string(dns, "remote_path").startswith("/"):
raise ValidationError("dns.remote_path must start with /")
if body_contains is not None and not isinstance(body_contains, str):
raise ValidationError("healthcheck.body_contains must be null or a string")
for float_key in ("timeout_seconds", "settle_seconds"):
value = health[float_key]
if isinstance(value, bool) or not isinstance(value, (int, float)) or not 0 <= value <= 120:
raise ValidationError(f"healthcheck.{float_key} must be in 0..120")
return Policy(
schema_version=1,
sing_box=SingBoxPolicy(
binary=_absolute(_string(sb, "binary"), "sing_box.binary"),
config_path=_absolute(_string(sb, "config_path"), "sing_box.config_path"),
service=service,
required_version=required_version,
),
runtime=RuntimePolicy(
uri_path=_absolute(_string(runtime, "uri_path"), "runtime.uri_path"),
state_dir=_absolute(_string(runtime, "state_dir"), "runtime.state_dir"),
lock_path=_absolute(_string(runtime, "lock_path"), "runtime.lock_path"),
backup_keep=_integer(runtime, "backup_keep", 1, 100),
),
network=NetworkPolicy(
upstream_interface=_interface(_string(network, "upstream_interface"), "upstream_interface"),
vpn_lan_interface=_interface(_string(network, "vpn_lan_interface"), "vpn_lan_interface"),
tun_name=_interface(_string(network, "tun_name"), "tun_name"),
tun_address=tun_address,
mtu=_integer(network, "mtu", 576, 9000),
route_exclude_address=tuple(exclusions),
iproute2_table_index=_integer(network, "iproute2_table_index", 1, 2**31 - 1),
iproute2_rule_index=_integer(network, "iproute2_rule_index", 1, 32765),
auto_redirect_input_mark=_mark(_string(network, "auto_redirect_input_mark"), "input mark"),
auto_redirect_output_mark=_mark(_string(network, "auto_redirect_output_mark"), "output mark"),
auto_redirect_reset_mark=_mark(_string(network, "auto_redirect_reset_mark"), "reset mark"),
auto_redirect_nfqueue=_integer(network, "auto_redirect_nfqueue", 0, 65535),
auto_redirect_fallback_rule_index=_integer(network, "auto_redirect_fallback_rule_index", 32766, 2**31 - 1),
),
dns=DnsPolicy(
bootstrap_server=bootstrap,
bootstrap_port=_integer(dns, "bootstrap_port", 1, 65535),
remote_server=remote,
remote_port=_integer(dns, "remote_port", 1, 65535),
remote_path=_string(dns, "remote_path"),
remote_tls_server_name=_string(dns, "remote_tls_server_name"),
strategy=_string(dns, "strategy"),
),
bandwidth=BandwidthPolicy(
up_mbps=_integer(bandwidth, "up_mbps", 0, 1_000_000),
down_mbps=_integer(bandwidth, "down_mbps", 0, 1_000_000),
),
healthcheck=HealthcheckPolicy(
url=url,
timeout_seconds=float(health["timeout_seconds"]),
settle_seconds=float(health["settle_seconds"]),
expected_status=_integer(health, "expected_status", 100, 599),
body_contains=body_contains,
),
)