500 lines
16 KiB
Python
500 lines
16 KiB
Python
#!/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()
|