271 lines
10 KiB
Python
271 lines
10 KiB
Python
# -*- encoding=utf8 -*-
|
||
import subprocess
|
||
import os
|
||
import shutil
|
||
import socket
|
||
import uuid
|
||
|
||
class DataManager:
|
||
"""负责文件的拉取及推送到局域网"""
|
||
DEFAULT_NET_USE_USER = "tplink"
|
||
DEFAULT_NET_USE_PASSWORD = "admin"
|
||
DEFAULT_NET_USE_TIMEOUT_SECONDS = 15
|
||
|
||
def __init__(self, config: dict):
|
||
self.config = config
|
||
self.worker_id = self._generate_worker_id()
|
||
self._authorized_shares = set()
|
||
print(f"[DataManager] Worker ID: {self.worker_id}")
|
||
|
||
def _get_ip_address(self):
|
||
"""获取本机IP地址(优先获取192.168开头的IP)"""
|
||
try:
|
||
hostname = socket.gethostname()
|
||
ip_addresses = socket.gethostbyname_ex(hostname)[2]
|
||
|
||
for ip in ip_addresses:
|
||
if ip.startswith('192.168'):
|
||
return ip
|
||
|
||
if ip_addresses:
|
||
return ip_addresses[0]
|
||
|
||
return '127.0.0.1'
|
||
except Exception as e:
|
||
print(f"[DataManager] 获取IP地址失败: {e}")
|
||
return '127.0.0.1'
|
||
|
||
def _get_mac_address(self) -> str:
|
||
"""获取本机Mac地址"""
|
||
try:
|
||
mac = uuid.getnode()
|
||
mac_address = ':'.join(['{:02x}'.format((mac >> elements) & 0xff) for elements in range(0, 8*6, 8)][::-1])
|
||
return mac_address.upper()
|
||
except Exception as e:
|
||
print(f"[DataManager] 获取MAC地址失败: {e}")
|
||
return '00:00:00:00:00:00'
|
||
|
||
def _generate_worker_id(self) -> str:
|
||
"""生成Worker ID (MAC_IP格式)"""
|
||
ip = self._get_ip_address()
|
||
mac = self._get_mac_address()
|
||
# 将MAC地址中的冒号替换为横线,使其适合作为目录名
|
||
mac_safe = mac.replace(':', '-')
|
||
return f"{mac_safe}_{ip}"
|
||
|
||
def _run_adb(self, args):
|
||
return subprocess.run(["adb"] + args, capture_output=True, text=True)
|
||
|
||
def _extract_unc_share_root(self, path):
|
||
"""从 UNC 路径提取共享根路径,如 \\host\share。"""
|
||
if not path:
|
||
return None
|
||
|
||
normalized = str(path).replace("/", "\\")
|
||
if not normalized.startswith("\\\\"):
|
||
return None
|
||
|
||
parts = [part for part in normalized.lstrip("\\").split("\\") if part]
|
||
if len(parts) < 2:
|
||
return None
|
||
return f"\\\\{parts[0]}\\{parts[1]}"
|
||
|
||
def _share_exists(self, share_root):
|
||
try:
|
||
return os.path.exists(share_root)
|
||
except OSError:
|
||
return False
|
||
|
||
def _ensure_share_access(self, path):
|
||
"""访问共享目录前,按需执行 net use 认证。"""
|
||
if os.name != "nt":
|
||
return True
|
||
|
||
share_root = self._extract_unc_share_root(path)
|
||
if not share_root:
|
||
return True
|
||
|
||
if share_root in self._authorized_shares and self._share_exists(share_root):
|
||
return True
|
||
|
||
if self._share_exists(share_root):
|
||
self._authorized_shares.add(share_root)
|
||
return True
|
||
|
||
print(f"[Data] 共享目录未就绪,尝试重新认证: {share_root}")
|
||
try:
|
||
subprocess.run(
|
||
["net", "use", share_root, "/delete", "/y"],
|
||
capture_output=True,
|
||
text=True,
|
||
timeout=self.DEFAULT_NET_USE_TIMEOUT_SECONDS,
|
||
)
|
||
res = subprocess.run(
|
||
[
|
||
"net",
|
||
"use",
|
||
share_root,
|
||
self.DEFAULT_NET_USE_PASSWORD,
|
||
f"/user:{self.DEFAULT_NET_USE_USER}",
|
||
],
|
||
capture_output=True,
|
||
text=True,
|
||
timeout=self.DEFAULT_NET_USE_TIMEOUT_SECONDS,
|
||
)
|
||
except Exception as e:
|
||
print(f"[Data] 执行 net use 失败: {e}")
|
||
return False
|
||
|
||
if res.returncode != 0:
|
||
error_msg = res.stderr.strip() or res.stdout.strip() or f"return code {res.returncode}"
|
||
print(f"[Data] 共享目录认证失败 {share_root}: {error_msg}")
|
||
return False
|
||
|
||
if not self._share_exists(share_root):
|
||
print(f"[Data] 共享目录认证后仍不可访问: {share_root}")
|
||
return False
|
||
|
||
self._authorized_shares.add(share_root)
|
||
print(f"[Data] 共享目录认证成功: {share_root}")
|
||
return True
|
||
|
||
def clear_remote_traffic_dir(self):
|
||
"""测试开始前清空设备上的流量文件夹"""
|
||
remote_dir = "/sdcard/Download/PCAPdroid"
|
||
print(f"[Data] 正在清空设备流量目录: {remote_dir}")
|
||
self._run_adb(["shell", f"rm -rf {remote_dir}/*"])
|
||
|
||
def sync_all_traffic_data(self):
|
||
"""全部测试完成后,统一从设备拉取整个流量文件夹并推送到局域网"""
|
||
print(f"\n[Data] 开始统一拉取所有流量数据...")
|
||
remote_dir = "/sdcard/Download/PCAPdroid"
|
||
|
||
# 1. 确保目标局域网目录存在 (使用 worker_id: MAC_IP)
|
||
target_share = os.path.join(self.config['TRAFFIC_DATA_SHARE'], self.worker_id)
|
||
if not self._ensure_share_access(target_share):
|
||
return False
|
||
|
||
try:
|
||
if not os.path.exists(target_share):
|
||
os.makedirs(target_share, exist_ok=True)
|
||
except Exception as e:
|
||
print(f"[Data] 无法创建局域网目录 {target_share}: {e}")
|
||
return False
|
||
|
||
# 2. 批量拉取到局域网目录
|
||
print(f"[Data] 正在将设备文件夹 {remote_dir} 拉取至 {target_share}")
|
||
res = self._run_adb(["pull", remote_dir, target_share])
|
||
|
||
if res.returncode == 0:
|
||
print(f"[Data] 流量数据同步成功!")
|
||
return True
|
||
else:
|
||
print(f"[Data] 流量数据同步失败: {res.stderr}")
|
||
return False
|
||
|
||
def sync_latest_traffic_file(self, package_name):
|
||
"""将设备上最新的流量 txt 文件拉取并推送到局域网指定目录"""
|
||
remote_traffic_dir = "/sdcard/Download/PCAPdroid/traffic"
|
||
|
||
# 1. 查找设备上最新的 txt 文件
|
||
# 使用 ls -t 按修改时间排序,取第一个
|
||
cmd = f"ls -t {remote_traffic_dir}/*.txt | head -n 1"
|
||
res = self._run_adb(["shell", cmd])
|
||
|
||
if res.returncode != 0 or not res.stdout.strip():
|
||
print(f"[Data] 未在设备上找到流量 txt 文件: {remote_traffic_dir}")
|
||
return False
|
||
|
||
remote_file_path = res.stdout.strip()
|
||
file_name = os.path.basename(remote_file_path)
|
||
|
||
# 1.5 确认文件内包含当前应用的流量数据(每行第一列为包名)
|
||
check_cmd = f"grep -c '^{package_name}' {remote_file_path} || true"
|
||
check_res = self._run_adb(["shell", check_cmd])
|
||
if check_res.returncode == 0 and check_res.stdout.strip().isdigit():
|
||
count = int(check_res.stdout.strip())
|
||
if count == 0:
|
||
print(f"[Data] 流量文件中未找到 {package_name} 的数据,跳过同步")
|
||
return False
|
||
else:
|
||
print(f"[Data] 无法检查流量文件内容,跳过同步")
|
||
return False
|
||
|
||
# 2. 准备局域网目标路径 (TRAFFIC_DATA_SHARE / worker_id(MAC_IP) / package_name)
|
||
target_dir = os.path.join(self.config['TRAFFIC_DATA_SHARE'], self.worker_id, package_name)
|
||
if not self._ensure_share_access(target_dir):
|
||
return False
|
||
|
||
try:
|
||
if not os.path.exists(target_dir):
|
||
os.makedirs(target_dir, exist_ok=True)
|
||
except Exception as e:
|
||
print(f"[Data] 无法创建局域网目录 {target_dir}: {e}")
|
||
return False
|
||
|
||
# 3. 拉取文件
|
||
target_file_path = os.path.join(target_dir, file_name)
|
||
print(f"[Data] 正在同步最新流量文件: {remote_file_path} -> {target_file_path}")
|
||
res = self._run_adb(["pull", remote_file_path, target_file_path])
|
||
|
||
if res.returncode == 0:
|
||
print(f"[Data] 单个流量文件同步成功!")
|
||
return True
|
||
else:
|
||
print(f"[Data] 单个流量文件同步失败: {res.stderr}")
|
||
return False
|
||
|
||
def sync_traversal_log(self, local_output_dir):
|
||
"""将 DroidBot 产生的本地 output 文件夹复制到局域网"""
|
||
if not local_output_dir or not os.path.exists(local_output_dir):
|
||
print("[Data] 未找到有效的输出文件夹,跳过同步")
|
||
return False
|
||
|
||
# 使用 worker_id (MAC_IP) 作为目录层级
|
||
folder_name = os.path.basename(local_output_dir)
|
||
target_parent = os.path.join(self.config['TRAVERSAL_LOG_SHARE'], self.worker_id)
|
||
target_dir = os.path.join(target_parent, folder_name)
|
||
if not self._ensure_share_access(target_parent):
|
||
return False
|
||
|
||
try:
|
||
if not os.path.exists(target_parent):
|
||
os.makedirs(target_parent, exist_ok=True)
|
||
|
||
print(f"[Data] 正在将日志复制至局域网: {target_dir}")
|
||
shutil.copytree(local_output_dir, target_dir, dirs_exist_ok=True)
|
||
return True
|
||
except Exception as e:
|
||
print(f"[Data] 复制日志文件夹失败: {e}")
|
||
return False
|
||
|
||
def sync_log_file(self, log_filepath, target_filename=None):
|
||
"""将 batch_run 的日志文件复制到局域网
|
||
|
||
Args:
|
||
log_filepath: 日志文件的完整路径
|
||
target_filename: 可选,指定目标文件名(不包含路径)
|
||
"""
|
||
if not log_filepath or not os.path.exists(log_filepath):
|
||
print("[Data] 未找到有效的日志文件,跳过同步")
|
||
return False
|
||
|
||
# 使用 worker_id (MAC_IP) 作为目录层级
|
||
file_name = target_filename if target_filename else os.path.basename(log_filepath)
|
||
target_dir = os.path.join(self.config.get('TRAVERSAL_LOG_SHARE'), self.worker_id)
|
||
target_file_path = os.path.join(target_dir, file_name)
|
||
if not self._ensure_share_access(target_dir):
|
||
return False
|
||
|
||
try:
|
||
if not os.path.exists(target_dir):
|
||
os.makedirs(target_dir, exist_ok=True)
|
||
|
||
print(f"[Data] 正在将日志文件复制至局域网: {target_file_path}")
|
||
shutil.copy2(log_filepath, target_file_path)
|
||
print(f"[Data] 日志文件同步成功!")
|
||
return True
|
||
except Exception as e:
|
||
print(f"[Data] 复制日志文件失败: {e}")
|
||
return False
|