autool-dispatcher/scripts/add_workers.py
2026-06-17 19:50:39 +08:00

500 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/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_targetIP用于去重
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()