from __future__ import annotations import json from collections import defaultdict from pathlib import Path from typing import TYPE_CHECKING from .connection import ConnectionFactory from .user import User if TYPE_CHECKING: from .connection import Connection class XrayConfig: def __init__(self, config_folder: Path, host: str): self.path = config_folder self.host = host self.inbounds_files = list(self.path.glob("*-in-*.json")) self.outbounds_files = list(self.path.glob("*-out-*.json")) self.inbounds_data = { f.name: self._load_json(f) for f in self.inbounds_files } self.outbounds_data = { f.name: self._load_json(f) for f in self.outbounds_files } def _load_json(self, path: Path) -> dict: return json.loads(path.read_text(encoding="utf-8")) def _save_json(self, data: dict, path: Path): path.write_text( json.dumps(data, indent=4, ensure_ascii=False) + "\n", encoding="utf-8", ) def _find_inbound(self, connection: Connection) -> tuple[str, dict]: for filename, data in self.inbounds_data.items(): inbounds = data.get("inbounds", []) for inbound in inbounds: if connection.matches_inbound(inbound): return filename, inbound raise RuntimeError(f"Inbound not found for {connection}") def find_managed_inbound_by_protocol(self, protocol: str) -> dict: for data in self.inbounds_data.values(): inbounds = data.get("inbounds", []) for inbound in inbounds: if not inbound.get("xrm"): continue if inbound.get("protocol") != protocol: continue return inbound raise RuntimeError( f"Managed inbound not found for protocol: {protocol}" ) def get_exit_points(self) -> set[str]: exit_points = set() for data in self.outbounds_data.values(): outbounds = data.get("outbounds", []) for outbound in outbounds: tag: str = outbound.get("tag") if tag: exit_point = tag.rsplit("-", 1)[1] exit_points.add(exit_point) exit_points.add("default") return exit_points def get_users(self) -> list[User]: connection_map: defaultdict[str, list[Connection]] = defaultdict(list) for data in self.inbounds_data.values(): inbounds = data.get("inbounds", []) for inbound in inbounds: protocol = inbound.get("protocol") if protocol in ConnectionFactory._registry: settings = inbound.get("settings", {}) clients = settings.get("clients", []) for client in clients: username, connection = ConnectionFactory.from_inbound( self.host, client, inbound ) connection_map[username].append(connection) users = [] for username, connections in connection_map.items(): users.append(User(username, connections)) return users def get_user(self, username: str) -> User | None: users = self.get_users() for user in users: if user.username == username: return user def add_user(self, user: User) -> None: modified_files = set() for connection in user.connections: filename, inbound = self._find_inbound(connection) connection.add_to_inbound(inbound, user.username) modified_files.add(filename) for filename in modified_files: path = self.path / filename self._save_json(self.inbounds_data[filename], path) def delete_user(self, username: str) -> bool: user = self.get_user(username) if not user: return False deleted = False modified_files = set() for connection in user.connections: try: filename, inbound = self._find_inbound(connection) except RuntimeError: continue connection.delete_from_inbound(inbound, username) modified_files.add(filename) deleted = True for filename in modified_files: path = self.path / filename self._save_json(self.inbounds_data[filename], path) return deleted def modify_user(self, user: User) -> bool: current_user = self.get_user(user.username) if current_user is None: return False old_connections = set(current_user.connections) new_connections = set(user.connections) to_add = new_connections - old_connections to_remove = old_connections - new_connections modified_files = set() for connection in to_remove: filename, inbound = self._find_inbound(connection) connection.delete_from_inbound(inbound, user.username) modified_files.add(filename) for connection in to_add: filename, inbound = self._find_inbound(connection) connection.add_to_inbound(inbound, user.username) modified_files.add(filename) for filename in modified_files: path = self.path / filename self._save_json(self.inbounds_data[filename], path) return True