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, ), )