Files
domainCheck/app/ui/domain_import.py
2026-04-14 22:53:52 +08:00

510 lines
20 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.
# -*- coding: UTF-8 -*-
'''
@Project :domainScanDemo
@File :domain_import.py
@IDE :PyCharm
@Author :梦伴
@Date :2026/4/8 23:47
@explain : 域名导入界面
'''
from PySide6.QtWidgets import QWidget, QVBoxLayout, QHBoxLayout, QPushButton, QLabel, QTextEdit, QFileDialog, QProgressBar
from PySide6.QtCore import Qt, QThread, Signal
from loguru import logger
from app.core.domain_collector import DomainCollector
class ImportThread(QThread):
"""
导入线程
"""
progress_updated = Signal(int)
finished = Signal(bool, str)
def __init__(self, domain_list, source_type):
"""
初始化导入线程
:param domain_list: 域名列表
:param source_type: 来源类型
"""
super().__init__()
self.domain_list = domain_list
self.source_type = source_type
def run(self):
"""
运行导入线程
"""
try:
collector = DomainCollector()
total = len(self.domain_list)
# 实时更新进度:开始
self.progress_updated.emit(0)
# 标准化域名和检查是否存在占30%进度)
normalized_domains = []
for i, domain in enumerate(self.domain_list):
from app.utils.domain_utils import normalize_domain
import tldextract
normalized = normalize_domain(domain)
if normalized:
# 提取顶级域名
ext = tldextract.extract(normalized)
tld = ext.suffix
normalized_domains.append((normalized, tld))
# 更新进度
progress = int((i + 1) / total * 30)
self.progress_updated.emit(progress)
# 批量检查域名是否存在
batch_data = []
existing_domains = []
if normalized_domains:
all_domains = [domain for domain, tld in normalized_domains]
existing_domains = collector.db.check_domains_exist(all_domains)
existing_set = set(existing_domains)
# 准备批量添加数据
for domain, tld in normalized_domains:
if domain not in existing_set:
batch_data.append((domain, tld, self.source_type))
# 分批次添加域名占70%进度)
batch_size = 1000
total_batches = len(batch_data)
for i in range(0, len(batch_data), batch_size):
batch = batch_data[i:i+batch_size]
collector.db.add_domains_batch(batch)
# 更新进度
processed = min(i + len(batch), total_batches)
progress = 30 + int(processed / total_batches * 70)
self.progress_updated.emit(progress)
# 完成导入
self.progress_updated.emit(100)
# 计算统计信息
stats = {
'total': total,
'valid': len(normalized_domains),
'added': len(batch_data),
'exists': len(existing_domains),
'invalid': total - len(normalized_domains),
'failed': 0
}
# 根据统计信息生成消息
if stats['added'] > 0:
message = f"导入完成: 总域名数 {stats['total']}, 有效域名数 {stats['valid']}, 新增域名数 {stats['added']}, 已存在域名数 {stats['exists']}, 无效域名数 {stats['invalid']}"
else:
message = f"导入完成: 所有域名已存在,未添加新域名"
self.finished.emit(True, message)
except Exception as e:
logger.error(f"导入失败: {e}")
self.finished.emit(False, f"导入失败: {str(e)}")
class ImportFileThread(QThread):
"""
文件导入线程,用于处理大文件
"""
progress_updated = Signal(int)
finished = Signal(bool, str)
def __init__(self, file_path, source_type):
"""
初始化文件导入线程
:param file_path: 文件路径
:param source_type: 来源类型
"""
super().__init__()
self.file_path = file_path
self.source_type = source_type
def run(self):
"""
运行文件导入线程
"""
try:
collector = DomainCollector()
# 首先计算文件中的域名数量
total = 0
encodings = ['utf-8', 'utf-8-sig', 'gbk', 'gb2312', 'cp936', 'latin-1', 'ascii']
encoding = 'utf-8' # 默认编码
# 尝试不同的编码格式计算域名数量
for enc in encodings:
try:
with open(self.file_path, 'r', encoding=enc) as f:
total = sum(1 for line in f if line.strip())
encoding = enc
break
except UnicodeDecodeError:
continue
if total == 0:
# 尝试使用二进制模式读取
try:
import chardet
with open(self.file_path, 'rb') as f:
raw_data = f.read()
result = chardet.detect(raw_data)
encoding = result['encoding']
if encoding:
total = sum(1 for line in raw_data.decode(encoding).split('\n') if line.strip())
else:
# 最后尝试使用 replace 模式读取
with open(self.file_path, 'r', encoding='utf-8', errors='replace') as f:
total = sum(1 for line in f if line.strip())
encoding = 'utf-8'
except Exception:
# 最后尝试使用 replace 模式读取
with open(self.file_path, 'r', encoding='utf-8', errors='replace') as f:
total = sum(1 for line in f if line.strip())
encoding = 'utf-8'
# 实时更新进度:开始
self.progress_updated.emit(0)
# 逐行读取文件并处理域名
normalized_domains = []
processed = 0
with open(self.file_path, 'r', encoding=encoding, errors='replace') as f:
for line in f:
domain = line.strip()
if domain:
from app.utils.domain_utils import normalize_domain
import tldextract
normalized = normalize_domain(domain)
if normalized:
# 提取顶级域名
ext = tldextract.extract(normalized)
tld = ext.suffix
normalized_domains.append((normalized, tld))
processed += 1
# 更新进度占30%
progress = int(processed / total * 30)
self.progress_updated.emit(progress)
# 批量检查域名是否存在
batch_data = []
existing_domains = []
if normalized_domains:
all_domains = [domain for domain, tld in normalized_domains]
existing_domains = collector.db.check_domains_exist(all_domains)
existing_set = set(existing_domains)
# 准备批量添加数据
for domain, tld in normalized_domains:
if domain not in existing_set:
batch_data.append((domain, tld, self.source_type))
# 分批次添加域名占70%进度)
batch_size = 1000
total_batches = len(batch_data)
for i in range(0, len(batch_data), batch_size):
batch = batch_data[i:i+batch_size]
collector.db.add_domains_batch(batch)
# 更新进度
processed_batches = min(i + len(batch), total_batches)
progress = 30 + int(processed_batches / total_batches * 70)
self.progress_updated.emit(progress)
# 完成导入
self.progress_updated.emit(100)
# 计算统计信息
stats = {
'total': total,
'valid': len(normalized_domains),
'added': len(batch_data),
'exists': len(existing_domains),
'invalid': total - len(normalized_domains),
'failed': 0
}
# 根据统计信息生成消息
if stats['added'] > 0:
message = f"导入完成: 总域名数 {stats['total']}, 有效域名数 {stats['valid']}, 新增域名数 {stats['added']}, 已存在域名数 {stats['exists']}, 无效域名数 {stats['invalid']}"
else:
message = f"导入完成: 所有域名已存在,未添加新域名"
self.finished.emit(True, message)
except Exception as e:
logger.error(f"导入失败: {e}")
self.finished.emit(False, f"导入失败: {str(e)}")
class DomainImportWidget(QWidget):
"""
域名导入界面
"""
def __init__(self):
"""
初始化域名导入界面
"""
super().__init__()
# 创建布局
layout = QVBoxLayout(self)
# 创建文本编辑框
self.text_edit = QTextEdit()
self.text_edit.setPlaceholderText("请输入域名,一行一个")
self.text_edit.setStyleSheet("""
QTextEdit {
font-size: 14px;
padding: 10px;
border: 1px solid #ddd;
border-radius: 4px;
background-color: #f9f9f9;
min-height: 300px;
}
""")
layout.addWidget(self.text_edit)
# 创建按钮布局
button_layout = QHBoxLayout()
# 导入文件按钮
self.import_file_btn = QPushButton("导入文件")
self.import_file_btn.clicked.connect(self.import_file)
self.import_file_btn.setStyleSheet("""
QPushButton {
font-size: 14px;
padding: 8px 16px;
background-color: #2196F3;
color: white;
border: none;
border-radius: 4px;
}
QPushButton:hover {
background-color: #0b7dda;
}
QPushButton:disabled {
background-color: #cccccc;
}
""")
button_layout.addWidget(self.import_file_btn)
# 开始导入按钮
self.start_import_btn = QPushButton("开始导入")
self.start_import_btn.clicked.connect(self.start_import)
self.start_import_btn.setStyleSheet("""
QPushButton {
font-size: 14px;
padding: 8px 16px;
background-color: #4CAF50;
color: white;
border: none;
border-radius: 4px;
margin-left: 10px;
}
QPushButton:hover {
background-color: #45a049;
}
QPushButton:disabled {
background-color: #cccccc;
}
""")
button_layout.addWidget(self.start_import_btn)
button_layout.setContentsMargins(0, 15, 0, 15)
layout.addLayout(button_layout)
# 创建进度条
self.progress_bar = QProgressBar()
self.progress_bar.setVisible(False)
self.progress_bar.setStyleSheet("""
QProgressBar {
height: 20px;
border: 1px solid #ddd;
border-radius: 10px;
background-color: #f0f0f0;
margin-bottom: 10px;
}
QProgressBar::chunk {
background-color: #4CAF50;
border-radius: 10px;
}
""")
layout.addWidget(self.progress_bar)
# 创建状态标签
self.status_label = QLabel("")
self.status_label.setAlignment(Qt.AlignCenter)
self.status_label.setStyleSheet("font-size: 14px; color: #333; padding: 10px; background-color: #f0f8ff; border-radius: 4px;")
layout.addWidget(self.status_label)
logger.info("域名导入界面创建完成")
def import_file(self):
"""
导入文件
"""
file_path, _ = QFileDialog.getOpenFileName(self, "选择文件", "", "文本文件 (*.txt)")
if file_path:
try:
# 尝试不同的编码格式
encodings = ['utf-8', 'utf-8-sig', 'gbk', 'gb2312', 'cp936', 'latin-1', 'ascii']
domain_count = 0
# 尝试使用不同编码读取并计数
for encoding in encodings:
try:
with open(file_path, 'r', encoding=encoding) as f:
domain_count = sum(1 for line in f if line.strip())
logger.info(f"使用编码 {encoding} 成功读取文件")
break
except UnicodeDecodeError:
continue
# 如果仍然失败,尝试使用二进制模式读取并猜测编码
if domain_count == 0:
try:
import chardet
with open(file_path, 'rb') as f:
raw_data = f.read()
result = chardet.detect(raw_data)
encoding = result['encoding']
if encoding:
domain_count = sum(1 for line in raw_data.decode(encoding).split('\n') if line.strip())
logger.info(f"使用 chardet 检测到编码 {encoding} 并成功读取文件")
else:
raise Exception("无法识别文件编码")
except Exception as e:
logger.warning(f"chardet 检测失败: {e}")
# 最后尝试使用 replace 模式读取
with open(file_path, 'r', encoding='utf-8', errors='replace') as f:
domain_count = sum(1 for line in f if line.strip())
logger.info("使用 utf-8 replace 模式读取文件")
# 对于大文件,不显示所有域名,只显示文件路径和域名数量
if domain_count > 1000:
self.text_edit.setText(f"文件路径: {file_path}\n域名数量: {domain_count}\n\n提示: 由于文件较大,仅显示文件信息,不显示具体域名。")
# 保存文件路径,用于后续导入
self.imported_file_path = file_path
else:
# 对于小文件,显示所有域名
domains = []
for encoding in encodings:
try:
with open(file_path, 'r', encoding=encoding) as f:
domains = f.readlines()
break
except UnicodeDecodeError:
continue
if not domains:
# 尝试使用二进制模式读取
try:
import chardet
with open(file_path, 'rb') as f:
raw_data = f.read()
result = chardet.detect(raw_data)
encoding = result['encoding']
if encoding:
domains = raw_data.decode(encoding).split('\n')
else:
with open(file_path, 'r', encoding='utf-8', errors='replace') as f:
domains = f.readlines()
except Exception:
with open(file_path, 'r', encoding='utf-8', errors='replace') as f:
domains = f.readlines()
domains = [domain.strip() for domain in domains if domain.strip()]
self.text_edit.setText('\n'.join(domains))
# 清除文件路径,使用文本框中的域名
self.imported_file_path = None
self.status_label.setText(f"成功读取 {domain_count} 个域名")
logger.info(f"成功读取文件: {file_path}, 共 {domain_count} 个域名")
except Exception as e:
self.status_label.setText(f"读取文件失败: {str(e)}")
logger.error(f"读取文件失败: {e}")
def start_import(self):
"""
开始导入
"""
# 检查是否有导入的文件路径
if hasattr(self, 'imported_file_path') and self.imported_file_path:
# 大文件导入,使用文件路径
file_path = self.imported_file_path
# 显示进度条
self.progress_bar.setVisible(True)
self.progress_bar.setValue(0)
self.status_label.setText("正在导入...")
# 禁用按钮
self.import_file_btn.setEnabled(False)
self.start_import_btn.setEnabled(False)
# 创建并启动导入线程
self.import_thread = ImportFileThread(file_path, 7) # 7 表示 TXT 导入
self.import_thread.progress_updated.connect(self.update_progress)
self.import_thread.finished.connect(self.import_finished)
self.import_thread.start()
logger.info(f"开始从文件导入: {file_path}")
else:
# 小文件或手动输入的域名
domains = self.text_edit.toPlainText().split('\n')
domains = [domain.strip() for domain in domains if domain.strip()]
if not domains:
self.status_label.setText("请输入域名")
return
# 显示进度条
self.progress_bar.setVisible(True)
self.progress_bar.setValue(0)
self.status_label.setText("正在导入...")
# 禁用按钮
self.import_file_btn.setEnabled(False)
self.start_import_btn.setEnabled(False)
# 创建并启动导入线程
self.import_thread = ImportThread(domains, 7) # 7 表示 TXT 导入
self.import_thread.progress_updated.connect(self.update_progress)
self.import_thread.finished.connect(self.import_finished)
self.import_thread.start()
logger.info(f"开始导入 {len(domains)} 个域名")
def update_progress(self, progress):
"""
更新进度
:param progress: 进度值
"""
self.progress_bar.setValue(progress)
def import_finished(self, success, message):
"""
导入完成
:param success: 是否成功
:param message: 消息
"""
self.status_label.setText(message)
self.progress_bar.setVisible(False)
# 启用按钮
self.import_file_btn.setEnabled(True)
self.start_import_btn.setEnabled(True)
logger.info(f"导入完成: {message}")