diff --git a/oneforall/collect.py b/oneforall/collect.py index 89cad7e..70e5279 100644 --- a/oneforall/collect.py +++ b/oneforall/collect.py @@ -75,7 +75,7 @@ class Collect(object): if not self.path: name = f'{self.domain}.{self.format}' self.path = config.result_save_path.joinpath(name) - dbexport.export(self.domain, dpath=self.path, format=self.format) + dbexport.export(self.domain, path=self.path, format=self.format) end = time.time() self.elapsed = round(end - start, 1) diff --git a/oneforall/common/utils.py b/oneforall/common/utils.py index b4c8388..c652a42 100644 --- a/oneforall/common/utils.py +++ b/oneforall/common/utils.py @@ -148,21 +148,28 @@ def get_semaphore(): def check_dpath(dpath=None): """ - 检查目录路径 + 检查结果输出目录路径 :param dpath: 传入的目录路径 :return: 目录路径 """ - if isinstance(dpath, str): - dpath = Path(dpath) + if dpath is None: + return config.result_save_path + try: + path = Path(dpath) + except Exception as e: + logger.log('ERROR', e.args) + path = config.result_save_path else: - dpath = config.result_save_path - if not dpath.is_dir(): - logger.log('FATAL', f'{dpath}不是目录') - if not dpath.exists(): - logger.log('ALERT', f'不存在{dpath}将会新建此目录') - dpath.mkdir(parents=True, exist_ok=True) - return dpath + if not path.is_dir(): + logger.log('ERROR', f'{path}不是目录') + path = config.result_save_path + if not path.exists(): + logger.log('ALERT', f'不存在{path}将会新建此目录') + path.mkdir(parents=True, exist_ok=True) + if path.resolve() == config.result_save_path: + logger.log('ALERT', f'使用默认结果输出目录{path}') + return path def check_format(format, count): @@ -187,17 +194,27 @@ def check_format(format, count): return 'csv' -def save_data(fpath, data): +def save_data(path, data): + """ + 保存结果数据到文件 + + :param path: 保存路径 + :param data: 待存数据 + :return: 保存成功与否 + """ try: - with open(fpath, 'w', encoding="utf-8", errors='ignore', newline='') as file: + with open(path, 'w', encoding="utf-8", errors='ignore', newline='') as file: file.write(data) - logger.log('ALERT', fpath) + logger.log('ALERT', f'结果输出{path}') + return True except TypeError: - with open(fpath, 'wb') as file: + with open(path, 'wb') as file: file.write(data) - logger.log('ALERT', fpath) + logger.log('ALERT', f'结果输出{path}') + return True except Exception as e: logger.log('ERROR', e.args) + return False def check_response(method, resp): diff --git a/oneforall/dbexport.py b/oneforall/dbexport.py index 8e9595f..ed3d65e 100644 --- a/oneforall/dbexport.py +++ b/oneforall/dbexport.py @@ -13,7 +13,7 @@ from common import utils from common.database import Database -def export(table, db=None, valid=None, dpath=None, format='csv', show=False): +def export(table, db=None, valid=None, path=None, format='csv', show=False): """ OneForAll数据库导出模块 @@ -31,13 +31,13 @@ def export(table, db=None, valid=None, dpath=None, format='csv', show=False): :param str db: 要导出的数据库路径(默认为results/result.sqlite3) :param int valid: 导出子域的有效性(默认None) :param str format: 导出格式(默认csv) - :param str dpath: 导出目录(默认None) + :param str path: 导出目录(默认None) :param bool show: 终端显示导出数据(默认False) """ - dpath = utils.check_dpath(dpath) + dir_path = utils.check_dpath(path) database = Database(db) - rows = database.export_data(table, valid) # 意外情况导出全部子域 + rows = database.export_data(table, valid) format = utils.check_format(format, len(rows)) if show: print(rows.dataset) @@ -46,8 +46,8 @@ def export(table, db=None, valid=None, dpath=None, format='csv', show=False): else: data = rows.export(format) database.close() - fpath = dpath.joinpath(f'{table}_subdomain.{format}') - utils.save_data(fpath, data) + file_path = dir_path.joinpath(f'{table}.{format}') + utils.save_data(file_path, data) if __name__ == '__main__': diff --git a/oneforall/oneforall.py b/oneforall/oneforall.py index c00b4f1..3a4b6b0 100644 --- a/oneforall/oneforall.py +++ b/oneforall/oneforall.py @@ -98,8 +98,8 @@ class OneForAll(object): self.dns = config.enable_dns_resolve if self.req is None: self.req = config.enable_http_request - old_table = self.domain + '_last' - new_table = self.domain + '_now' + old_table = self.domain + '_last_result' + new_table = self.domain + '_now_result' collect = Collect(self.domain, export=False) collect.run() if self.brute: @@ -108,7 +108,8 @@ class OneForAll(object): brute.run() db = Database() - db.copy_table(self.domain, self.domain+'_ori') + original_table = self.domain + '_original_result' + db.copy_table(self.domain, original_table) db.remove_invalid(self.domain) db.deduplicate_subdomain(self.domain) @@ -122,7 +123,6 @@ class OneForAll(object): # 不解析子域直接导出结果 if not self.dns: # 数据库导出 - self.valid = None dbexport.export(self.domain, valid=self.valid, format=self.format, show=self.show) db.drop_table(new_table) @@ -135,6 +135,7 @@ class OneForAll(object): # 标记新发现子域 self.data = utils.mark_subdomain(old_data, self.data) + # 获取事件循环 loop = asyncio.get_event_loop() asyncio.set_event_loop(loop) @@ -143,7 +144,7 @@ class OneForAll(object): self.data = loop.run_until_complete(task) # 保存解析结果 - resolve_table = self.domain + '_res' + resolve_table = self.domain + '_resolve_result' db.drop_table(resolve_table) db.create_table(resolve_table) db.save_db(resolve_table, self.data, 'resolve') @@ -151,8 +152,7 @@ class OneForAll(object): # 不请求子域直接导出结果 if not self.req: # 数据库导出 - self.valid = None - dbexport.export(self.domain, valid=self.valid, + dbexport.export(resolve_table, valid=self.valid, format=self.format, show=self.show) db.drop_table(new_table) db.rename_table(self.domain, new_table) diff --git a/oneforall/takeover.py b/oneforall/takeover.py index bf5dc15..035726f 100644 --- a/oneforall/takeover.py +++ b/oneforall/takeover.py @@ -56,21 +56,21 @@ class Takeover(Module): Note: 参数format可选格式有'txt', 'rst', 'csv', 'tsv', 'json', 'yaml', 'html', 'jira', 'xls', 'xlsx', 'dbf', 'latex', 'ods' - 参数dpath为None默认使用OneForAll结果目录 + 参数path为None默认使用OneForAll结果目录 :param any target: 单个子域或者每行一个子域的文件路径(必需参数) :param int thread: 线程数(默认100) :param str format: 导出格式(默认csv) - :param str dpath: 导出目录(默认None) + :param str path: 导出目录(默认None) """ - def __init__(self, target, thread=100, dpath=None, format='csv'): + def __init__(self, target, thread=100, path=None, format='csv'): Module.__init__(self) self.subdomains = set() self.module = 'Check' self.source = 'Takeover' self.target = target self.thread = thread - self.dpath = dpath + self.path = path self.format = format self.fingerprints = None self.subdomainq = Queue() @@ -78,14 +78,14 @@ class Takeover(Module): self.results = Dataset() def save(self): - logger.log('INFOR', '正在保存检查结果') + logger.log('DEBUG', '正在保存检查结果') if self.format == 'txt': data = str(self.results) else: data = self.results.export(self.format) timestamp = utils.get_timestamp() - fpath = self.dpath.joinpath(f'takeover_{timestamp}.{self.format}') - utils.save_data(fpath, data) + path = self.path.joinpath(f'takeover_{timestamp}.{self.format}') + utils.save_data(path, data) def compare(self, subdomain, cname, responses): domain_resp = self.get('http://' + subdomain, check=False) @@ -136,7 +136,7 @@ class Takeover(Module): logger.log('INFOR', f'开始执行{self.source}模块') self.subdomains = utils.get_domains(self.target) self.format = utils.check_format(self.format, len(self.subdomains)) - self.dpath = utils.check_dpath(self.dpath) + self.path = utils.check_dpath(self.path) if self.subdomains: logger.log('INFOR', f'正在检查子域接管风险') self.fingerprints = get_fingerprint()