#!/usr/bin/env python3 # -*- encoding=utf8 -*- """ Worker 批量添加脚本 扫描指定 IP 范围内的设备,通过 SSH 获取 MAC 地址, 自动生成 Worker 条目并合并写入 config/config.json 的 WORKER_INVENTORY。 使用示例: # 扫描 192.168.1.10 到 192.168.1.50 python scripts/add_workers.py --range 192.168.1.10-50 # 扫描多个范围 python scripts/add_workers.py --range 192.168.1.10-50 192.168.2.50-80 # 指定标签 python scripts/add_workers.py --range 192.168.1.10-50 --tags group_b # 仅扫描,不写入 python scripts/add_workers.py --range 192.168.1.10-50 --dry-run # 跳过连通性检测 python scripts/add_workers.py --range 192.168.1.10-50 --skip-ping """ import argparse import json import os import re import socket import sys from concurrent.futures import ThreadPoolExecutor, as_completed from typing import Dict, List, Optional, Tuple # 项目根目录 BASE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) sys.path.insert(0, BASE_DIR) from config import ( CONFIG_PATH, SSH_DEFAULT_USER, SSH_DEFAULT_PASSWORD, SSH_DEFAULT_PORT, WORKER_INVENTORY_CONFIG_KEY, WORKER_INVENTORY_PATH, ) # 默认 Worker 配置 DEFAULT_REPO_DIR = "D:/autool" DEFAULT_PYTHON_EXE = "python" # 并发设置 MAX_PING_WORKERS = 50 MAX_SSH_WORKERS = 10 PING_TIMEOUT = 2 # 秒 SSH_TIMEOUT = 10 # 秒 def parse_ip_range(range_str: str) -> List[str]: """ 解析 IP 范围字符串,支持以下格式: - 单个 IP: 192.168.1.10 - 末段范围: 192.168.1.10-50 - CIDR: 192.168.1.0/24 """ range_str = range_str.strip() # CIDR 格式: 192.168.1.0/24 cidr_match = re.fullmatch(r"(\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3})/(\d{1,2})", range_str) if cidr_match: base_ip = cidr_match.group(1) prefix_len = int(cidr_match.group(2)) if prefix_len < 24 or prefix_len > 32: raise ValueError(f"仅支持 /24 ~ /32 的 CIDR 范围: {range_str}") parts = base_ip.split(".") base = int(parts[3]) host_bits = 32 - prefix_len count = 2 ** host_bits # 排除网络地址和广播地址 start = max(base, 1) end = min(base + count - 1, 254) return [f"{parts[0]}.{parts[1]}.{parts[2]}.{i}" for i in range(start, end + 1)] # 末段范围格式: 192.168.1.10-50 range_match = re.fullmatch(r"(\d{1,3}\.\d{1,3}\.\d{1,3})\.(\d{1,3})-(\d{1,3})", range_str) if range_match: prefix = range_match.group(1) start = int(range_match.group(2)) end = int(range_match.group(3)) if start > end: raise ValueError(f"起始地址大于结束地址: {range_str}") if end > 254: raise ValueError(f"IP 末段超出范围 (最大 254): {range_str}") return [f"{prefix}.{i}" for i in range(start, end + 1)] # 单个 IP: 192.168.1.10 ip_match = re.fullmatch(r"\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}", range_str) if ip_match: return [range_str] raise ValueError(f"无法解析 IP 范围格式: {range_str}") def check_ssh_port(ip: str, port: int = SSH_DEFAULT_PORT, timeout: float = PING_TIMEOUT) -> bool: """尝试通过 TCP 连接 SSH 端口检查主机是否在线""" try: with socket.create_connection((ip, port), timeout=timeout): return True except (socket.timeout, ConnectionRefusedError, OSError): return False def scan_reachable_hosts(ip_list: List[str], port: int = SSH_DEFAULT_PORT, skip_ping: bool = False) -> List[str]: """并发端口扫描,返回 SSH 端口可达的主机列表""" if skip_ping: print(f" 跳过连通性检测,将直接尝试 SSH 登录 {len(ip_list)} 台设备") return ip_list print(f" 正在扫描 {len(ip_list)} 个 IP 地址的 SSH 端口...") reachable = [] with ThreadPoolExecutor(max_workers=MAX_PING_WORKERS) as executor: futures = {executor.submit(check_ssh_port, ip, port): ip for ip in ip_list} for future in as_completed(futures): ip = futures[future] try: if future.result(): reachable.append(ip) print(f" ✓ {ip} 在线 (端口可达)") else: print(f" ✗ {ip} 不可达 (端口关闭或超时)") except Exception as e: print(f" ✗ {ip} 检测异常: {e}") # 按 IP 排序 reachable.sort(key=lambda x: tuple(int(p) for p in x.split("."))) print(f" 扫描完成,{len(reachable)} 台设备 SSH 端口可连接\n") return reachable def get_mac_via_ssh( ip: str, user: str = SSH_DEFAULT_USER, password: str = SSH_DEFAULT_PASSWORD, port: int = SSH_DEFAULT_PORT, ) -> Optional[str]: """ 通过 SSH 连接 Windows 设备并获取物理网卡 MAC 地址。 使用 getmac /fo csv /nh 获取 MAC 列表,选取以太网适配器的 MAC。 """ try: import paramiko except ImportError: print(" [错误] 需要 paramiko 库,请运行: pip install paramiko") sys.exit(1) ssh = paramiko.SSHClient() ssh.set_missing_host_key_policy(paramiko.AutoAddPolicy()) try: ssh.connect( hostname=ip, port=port, username=user, password=password, timeout=SSH_TIMEOUT, auth_timeout=SSH_TIMEOUT, banner_timeout=SSH_TIMEOUT, look_for_keys=True, allow_agent=True, ) # 获取 MAC 地址列表 _, stdout, _ = ssh.exec_command("getmac /fo csv /nh", timeout=SSH_TIMEOUT) output = stdout.read().decode("utf-8", errors="ignore").strip() if not output: return None # 解析 getmac 输出,格式: "MAC地址","传输名称","..." # 优先选取以太网/有线网卡的 MAC,跳过虚拟网卡和断开的连接 best_mac = None for line in output.splitlines(): line = line.strip() if not line: continue # 去掉引号分割 parts = [p.strip().strip('"') for p in line.split(",")] if len(parts) < 2: continue mac = parts[0].strip() transport = parts[1].strip() if len(parts) > 1 else "" # 跳过无效 MAC if not re.match(r"([0-9A-Fa-f]{2}[:-]){5}[0-9A-Fa-f]{2}", mac): continue # 跳过已断开的连接 if "已断开" in transport or "Disconnected" in transport.lower() or "Media disconnected" in transport.lower(): continue # 统一 MAC 格式为 XX:XX:XX:XX:XX:XX(大写、冒号分隔) mac = mac.upper().replace("-", ":") # 优先选取硬件以太网适配器 transport_lower = transport.lower() if any(kw in transport_lower for kw in ["ethernet", "以太网", "realtek", "intel"]): best_mac = mac break # 记录第一个有效的 MAC 作为备选 if best_mac is None: best_mac = mac return best_mac except Exception as e: print(f" ✗ {ip} SSH 连接失败: {e}") return None finally: ssh.close() def collect_worker_info( ip_list: List[str], tags: List[str], user: str = SSH_DEFAULT_USER, password: str = SSH_DEFAULT_PASSWORD, port: int = SSH_DEFAULT_PORT, ) -> List[Dict]: """批量通过 SSH 收集设备信息,生成 worker 条目""" print(f" 正在通过 SSH 获取 {len(ip_list)} 台设备的 MAC 地址...") workers = [] def _collect_one(ip: str) -> Optional[Dict]: mac = get_mac_via_ssh(ip, user=user, password=password, port=port) if mac: worker_id = f"{ip}_{mac}" print(f" ✓ {ip} → MAC: {mac} → worker_id: {worker_id}") return { "worker_id": worker_id, "ssh_target": ip, "repo_dir": DEFAULT_REPO_DIR, "python_exe": DEFAULT_PYTHON_EXE, "tags": list(tags), } else: print(f" ✗ {ip} 无法获取 MAC 地址") return None with ThreadPoolExecutor(max_workers=MAX_SSH_WORKERS) as executor: futures = {executor.submit(_collect_one, ip): ip for ip in ip_list} for future in as_completed(futures): result = future.result() if result: workers.append(result) # 按 IP 排序 workers.sort(key=lambda w: tuple(int(p) for p in w["ssh_target"].split("."))) print(f" 成功获取 {len(workers)} 台设备信息\n") return workers def load_existing_inventory(path: str) -> Tuple[List[Dict], set, set]: """加载现有 Worker 列表,返回列表、已存在 worker_id 和已存在 IP 集合""" if not os.path.exists(path): return [], set(), set() with open(path, "r", encoding="utf-8") as f: payload = json.load(f) data = payload.get(WORKER_INVENTORY_CONFIG_KEY, []) if isinstance(payload, dict) else payload if not isinstance(data, list): print(f" [警告] {path} 格式异常(非数组),将创建新文件") return [], set(), set() existing_ids = set() # 同时收集已存在的 ssh_target(IP)用于去重 existing_ips = set() for item in data: if isinstance(item, dict): wid = item.get("worker_id", "") if wid: existing_ids.add(wid) ip = item.get("ssh_target", "") if ip: existing_ips.add(ip) return data, existing_ids, existing_ips def save_inventory(path: str, workers: List[Dict]) -> None: """写入统一配置;如果传入旧数组文件路径,则按旧格式写回。""" payload = {} if os.path.exists(path): with open(path, "r", encoding="utf-8") as f: try: payload = json.load(f) except json.JSONDecodeError: payload = {} if isinstance(payload, dict): payload[WORKER_INVENTORY_CONFIG_KEY] = workers with open(path, "w", encoding="utf-8") as f: json.dump(payload, f, indent=2, ensure_ascii=False) f.write("\n") return with open(path, "w", encoding="utf-8") as f: json.dump(workers, f, indent=2, ensure_ascii=False) f.write("\n") def merge_and_save( existing: List[Dict], existing_ids: set, existing_ips: set, new_workers: List[Dict], path: str, dry_run: bool = False, ) -> int: """合并新旧 Worker 列表并写入文件""" to_add = [] skipped_dup_id = 0 skipped_dup_ip = 0 for worker in new_workers: wid = worker["worker_id"] ip = worker["ssh_target"] if wid in existing_ids: print(f" [跳过] worker_id 已存在: {wid}") skipped_dup_id += 1 continue if ip in existing_ips: print(f" [跳过] IP 已存在于其他 worker: {ip}") skipped_dup_ip += 1 continue to_add.append(worker) if not to_add: print(" 没有新的 Worker 需要添加") return 0 print(f"\n 将添加 {len(to_add)} 台新 Worker:") for w in to_add: tag_str = f" tags={w['tags']}" if w["tags"] else "" print(f" + {w['worker_id']}{tag_str}") if skipped_dup_id > 0: print(f" (已跳过 {skipped_dup_id} 台 worker_id 重复的设备)") if skipped_dup_ip > 0: print(f" (已跳过 {skipped_dup_ip} 台 IP 重复的设备)") if dry_run: print("\n [Dry Run] 未写入文件") return len(to_add) # 合并并写入 merged = existing + to_add # 备份原文件 if os.path.exists(path): backup_path = path + ".bak" with open(path, "r", encoding="utf-8") as f: backup_content = f.read() with open(backup_path, "w", encoding="utf-8") as f: f.write(backup_content) print(f"\n 已备份原文件到: {backup_path}") save_inventory(path, merged) print(f" 已写入 {path},共 {len(merged)} 台 Worker") return len(to_add) def main(): parser = argparse.ArgumentParser( description="Worker 批量添加工具 — 扫描 IP 范围并自动添加到 config/config.json", formatter_class=argparse.RawDescriptionHelpFormatter, epilog=""" 示例: python scripts/add_workers.py --range 192.168.1.10-50 python scripts/add_workers.py --range 192.168.1.10-50 192.168.2.50-80 python scripts/add_workers.py --range 192.168.1.10-50 --tags group_b python scripts/add_workers.py --range 192.168.1.10-50 --dry-run python scripts/add_workers.py --range 192.168.1.10-50 --skip-ping """, ) parser.add_argument( "--range", nargs="+", required=True, dest="ip_ranges", help="IP 范围,支持格式: 192.168.1.10 / 192.168.1.10-50 / 192.168.1.0/24", ) parser.add_argument( "--tags", nargs="*", default=[], help="为新添加的 Worker 设置标签", ) parser.add_argument( "--dry-run", action="store_true", help="仅扫描和预览,不写入文件", ) parser.add_argument( "--skip-ping", action="store_true", help="跳过连通性检测,直接尝试 SSH 登录", ) parser.add_argument( "--inventory", default=WORKER_INVENTORY_PATH, help=f"配置文件或旧 worker_inventory.json 路径 (默认: {CONFIG_PATH})", ) parser.add_argument( "--user", default=SSH_DEFAULT_USER, help=f"SSH 用户名 (默认: {SSH_DEFAULT_USER})", ) parser.add_argument( "--password", default=SSH_DEFAULT_PASSWORD, help="SSH 密码", ) parser.add_argument( "--port", type=int, default=SSH_DEFAULT_PORT, help=f"SSH 端口 (默认: {SSH_DEFAULT_PORT})", ) args = parser.parse_args() print("=" * 60) print("Worker 批量添加工具") print("=" * 60) # 1. 解析 IP 范围 all_ips = [] for r in args.ip_ranges: try: ips = parse_ip_range(r) print(f" 范围 {r} → {len(ips)} 个 IP") all_ips.extend(ips) except ValueError as e: print(f" [错误] {e}") sys.exit(1) # 去重并排序 all_ips = sorted(set(all_ips), key=lambda x: tuple(int(p) for p in x.split("."))) print(f"\n 共计 {len(all_ips)} 个唯一 IP 地址\n") if not all_ips: print(" 没有有效的 IP 地址") return # 2. 端口扫描 print("[1/4] 连通性检测") reachable = scan_reachable_hosts(all_ips, port=args.port, skip_ping=args.skip_ping) if not reachable: print(" 没有在线的设备") return # 3. SSH 获取 MAC 地址 print("[2/4] 获取设备信息") new_workers = collect_worker_info( reachable, tags=args.tags, user=args.user, password=args.password, port=args.port, ) if not new_workers: print(" 未能获取任何设备信息") return # 4. 加载现有 inventory print("[3/4] 加载现有 Worker 列表") existing, existing_ids, existing_ips = load_existing_inventory(args.inventory) print(f" 当前共有 {len(existing)} 台 Worker\n") # 5. 合并写入 print("[4/4] 合并并写入") added = merge_and_save( existing, existing_ids, existing_ips, new_workers, args.inventory, dry_run=args.dry_run, ) print("\n" + "=" * 60) print(f"完成!新增 {added} 台 Worker") if args.dry_run: print("(Dry Run 模式,未实际写入)") print("=" * 60) if __name__ == "__main__": main()