实现配置文件插拔式设计

This commit is contained in:
Jing Ling
2020-08-07 20:06:16 +08:00
parent d7f7c54b1a
commit 69ce692ff1
14 changed files with 103 additions and 103 deletions
+22 -22
View File
@@ -21,7 +21,7 @@ from dns.resolver import NXDOMAIN, YXDOMAIN, NoAnswer, NoNameservers
import dbexport import dbexport
from common import utils from common import utils
from config import setting from config import settings
from common.module import Module from common.module import Module
from config.log import logger from config.log import logger
@@ -241,21 +241,21 @@ def collect_wildcard_record(domain, authoritative_ns):
def get_nameservers_path(enable_wildcard, ns_ip_list): def get_nameservers_path(enable_wildcard, ns_ip_list):
path = setting.brute_nameservers_path path = settings.brute_nameservers_path
if not enable_wildcard: if not enable_wildcard:
return path return path
if not ns_ip_list: if not ns_ip_list:
return path return path
path = setting.authoritative_dns_path path = settings.authoritative_dns_path
ns_data = '\n'.join(ns_ip_list) ns_data = '\n'.join(ns_ip_list)
utils.save_data(path, ns_data) utils.save_data(path, ns_data)
return path return path
def check_dict(): def check_dict():
if not setting.enable_check_dict: if not settings.enable_check_dict:
return return
sec = setting.check_time sec = settings.check_time
logger.log('ALERT', f'You have {sec} seconds to check ' logger.log('ALERT', f'You have {sec} seconds to check '
f'whether the configuration is correct or not') f'whether the configuration is correct or not')
logger.log('ALERT', f'If you want to exit, please use `Ctrl + C`') logger.log('ALERT', f'If you want to exit, please use `Ctrl + C`')
@@ -401,13 +401,13 @@ def check_ip_times(times):
:param times: IP address times :param times: IP address times
:return bool: result :return bool: result
""" """
if times > setting.ip_appear_maximum: if times > settings.ip_appear_maximum:
return True return True
return False return False
def is_valid_subdomain(ip, ttl, times, wc_ips, wc_ttl): def is_valid_subdomain(ip, ttl, times, wc_ips, wc_ttl):
ip_blacklist = setting.brute_ip_blacklist ip_blacklist = settings.brute_ip_blacklist
if ip in ip_blacklist: # 解析ip在黑名单ip则为非法子域 if ip in ip_blacklist: # 解析ip在黑名单ip则为非法子域
return 0, 'IP blacklist' return 0, 'IP blacklist'
if all([wc_ips, wc_ttl]): # 有泛解析记录才进行对比 if all([wc_ips, wc_ttl]): # 有泛解析记录才进行对比
@@ -426,9 +426,9 @@ def save_brute_dict(dict_path, dict_set):
def delete_file(dict_path, output_paths): def delete_file(dict_path, output_paths):
if setting.delete_generated_dict: if settings.delete_generated_dict:
dict_path.unlink() dict_path.unlink()
if setting.delete_massdns_result: if settings.delete_massdns_result:
for output_path in output_paths: for output_path in output_paths:
output_path.unlink() output_path.unlink()
@@ -478,15 +478,15 @@ class Brute(Module):
self.source = 'Brute' self.source = 'Brute'
self.target = target self.target = target
self.process_num = process or utils.get_process_num() self.process_num = process or utils.get_process_num()
self.concurrent_num = concurrent or setting.brute_concurrent_num self.concurrent_num = concurrent or settings.brute_concurrent_num
self.word = word self.word = word
self.wordlist = wordlist or setting.brute_wordlist_path self.wordlist = wordlist or settings.brute_wordlist_path
self.recursive_brute = recursive or setting.enable_recursive_brute self.recursive_brute = recursive or settings.enable_recursive_brute
self.recursive_depth = depth or setting.brute_recursive_depth self.recursive_depth = depth or settings.brute_recursive_depth
self.recursive_nextlist = nextlist or setting.recursive_nextlist_path self.recursive_nextlist = nextlist or settings.recursive_nextlist_path
self.fuzz = fuzz or setting.enable_fuzz self.fuzz = fuzz or settings.enable_fuzz
self.place = place or setting.fuzz_place self.place = place or settings.fuzz_place
self.rule = rule or setting.fuzz_rule self.rule = rule or settings.fuzz_rule
self.export = export self.export = export
self.alive = alive self.alive = alive
self.format = format self.format = format
@@ -496,8 +496,8 @@ class Brute(Module):
self.domain = str() # 当前正在进行爆破的域名 self.domain = str() # 当前正在进行爆破的域名
self.ips_times = dict() # IP集合出现次数 self.ips_times = dict() # IP集合出现次数
self.enable_wildcard = False # 当前域名是否使用泛解析 self.enable_wildcard = False # 当前域名是否使用泛解析
self.wildcard_check = setting.brute_wildcard_check self.wildcard_check = settings.brute_wildcard_check
self.wildcard_deal = setting.brute_wildcard_deal self.wildcard_deal = settings.brute_wildcard_deal
self.check_env = True self.check_env = True
self.quite = False self.quite = False
@@ -557,8 +557,8 @@ class Brute(Module):
def main(self, domain): def main(self, domain):
start = time.time() start = time.time()
logger.log('INFOR', f'Blasting {domain} ') logger.log('INFOR', f'Blasting {domain} ')
massdns_dir = setting.third_party_dir.joinpath('massdns') massdns_dir = settings.third_party_dir.joinpath('massdns')
result_dir = setting.result_save_dir result_dir = settings.result_save_dir
temp_dir = result_dir.joinpath('temp') temp_dir = result_dir.joinpath('temp')
utils.check_dir(temp_dir) utils.check_dir(temp_dir)
massdns_path = utils.get_massdns_path(massdns_dir) massdns_path = utils.get_massdns_path(massdns_dir)
@@ -646,7 +646,7 @@ class Brute(Module):
logger.log('INFOR', f'Finished {self.source} module\'s brute {self.domain}') logger.log('INFOR', f'Finished {self.source} module\'s brute {self.domain}')
if not self.path: if not self.path:
name = f'{self.domain}_brute_result.{self.format}' name = f'{self.domain}_brute_result.{self.format}'
self.path = setting.result_save_dir.joinpath(name) self.path = settings.result_save_dir.joinpath(name)
# 数据库导出 # 数据库导出
if self.export: if self.export:
dbexport.export(self.domain, dbexport.export(self.domain,
+2 -2
View File
@@ -9,7 +9,7 @@ import records
from records import Connection from records import Connection
from config.log import logger from config.log import logger
from config import setting from config import settings
class Database(object): class Database(object):
@@ -29,7 +29,7 @@ class Database(object):
return db_path return db_path
protocol = 'sqlite:///' protocol = 'sqlite:///'
if not db_path: # 数据库路径为空连接默认数据库 if not db_path: # 数据库路径为空连接默认数据库
db_path = f'{protocol}{setting.result_save_dir}/result.sqlite3' db_path = f'{protocol}{settings.result_save_dir}/result.sqlite3'
else: else:
db_path = protocol + db_path db_path = protocol + db_path
db = records.Database(db_path) # 不存在数据库时会新建一个数据库 db = records.Database(db_path) # 不存在数据库时会新建一个数据库
+2 -2
View File
@@ -1,6 +1,6 @@
import re import re
import tldextract import tldextract
from config import setting from config import settings
class Domain(object): class Domain(object):
@@ -38,7 +38,7 @@ class Domain(object):
:return: extracted domain results :return: extracted domain results
""" """
data_storage_dir = setting.data_storage_dir data_storage_dir = settings.data_storage_dir
extract_cache_file = data_storage_dir.joinpath('public_suffix_list.dat') extract_cache_file = data_storage_dir.joinpath('public_suffix_list.dat')
tldext = tldextract.TLDExtract(extract_cache_file, None) tldext = tldextract.TLDExtract(extract_cache_file, None)
result = self.match() result = self.match()
+10 -10
View File
@@ -9,7 +9,7 @@ import time
import requests import requests
from config.log import logger from config.log import logger
from config import setting from config import settings
from common import utils from common import utils
from common.domain import Domain from common.domain import Domain
from common.database import Database from common.database import Database
@@ -24,9 +24,9 @@ class Module(object):
self.cookie = None self.cookie = None
self.header = dict() self.header = dict()
self.proxy = None self.proxy = None
self.delay = setting.request_delay # 请求睡眠时延 self.delay = settings.request_delay # 请求睡眠时延
self.timeout = setting.request_timeout # 请求超时时间 self.timeout = settings.request_timeout # 请求超时时间
self.verify = setting.request_verify # 请求SSL验证 self.verify = settings.request_verify # 请求SSL验证
self.domain = str() # 当前进行子域名收集的主域 self.domain = str() # 当前进行子域名收集的主域
self.type = 'A' # 对主域进行子域收集时利用的DNS记录查询类型(默认利用A记录) self.type = 'A' # 对主域进行子域收集时利用的DNS记录查询类型(默认利用A记录)
self.subdomains = set() # 存放发现的子域 self.subdomains = set() # 存放发现的子域
@@ -185,7 +185,7 @@ class Module(object):
:return: header :return: header
""" """
# logger.log('DEBUG', f'Get request header') # logger.log('DEBUG', f'Get request header')
if setting.enable_fake_header: if settings.enable_fake_header:
return utils.gen_fake_header() return utils.gen_fake_header()
else: else:
return self.header return self.header
@@ -197,13 +197,13 @@ class Module(object):
:param str module: module name :param str module: module name
:return: proxy :return: proxy
""" """
if not setting.enable_proxy: if not settings.enable_proxy:
logger.log('TRACE', f'All modules do not use proxy') logger.log('TRACE', f'All modules do not use proxy')
return self.proxy return self.proxy
if setting.proxy_all_module: if settings.proxy_all_module:
logger.log('TRACE', f'{module} module uses proxy') logger.log('TRACE', f'{module} module uses proxy')
return utils.get_random_proxy() return utils.get_random_proxy()
if module in setting.proxy_partial_module: if module in settings.proxy_partial_module:
logger.log('TRACE', f'{module} module uses proxy') logger.log('TRACE', f'{module} module uses proxy')
return utils.get_random_proxy() return utils.get_random_proxy()
else: else:
@@ -219,11 +219,11 @@ class Module(object):
:return bool: whether saved successfully :return bool: whether saved successfully
""" """
if not setting.save_module_result: if not settings.save_module_result:
return False return False
logger.log('TRACE', f'Save the subdomain results found by ' logger.log('TRACE', f'Save the subdomain results found by '
f'{self.source} module as a json file') f'{self.source} module as a json file')
path = setting.result_save_dir.joinpath(self.domain, self.module) path = settings.result_save_dir.joinpath(self.domain, self.module)
path.mkdir(parents=True, exist_ok=True) path.mkdir(parents=True, exist_ok=True)
name = self.source + '.json' name = self.source + '.json'
path = path.joinpath(name) path = path.joinpath(name)
+14 -14
View File
@@ -9,11 +9,11 @@ from bs4 import BeautifulSoup
from common import utils from common import utils
from config.log import logger from config.log import logger
from config import setting from config import settings
def get_limit_conn(): def get_limit_conn():
limit_open_conn = setting.limit_open_conn limit_open_conn = settings.limit_open_conn
if limit_open_conn is None: # 默认情况 if limit_open_conn is None: # 默认情况
limit_open_conn = utils.get_semaphore() limit_open_conn = utils.get_semaphore()
elif not isinstance(limit_open_conn, int): # 如果传入不是数字的情况 elif not isinstance(limit_open_conn, int): # 如果传入不是数字的情况
@@ -31,7 +31,7 @@ def get_ports(port):
ports = {port} ports = {port}
elif port in {'small', 'medium', 'large'}: elif port in {'small', 'medium', 'large'}:
logger.log('DEBUG', f'{port} port range') logger.log('DEBUG', f'{port} port range')
ports = setting.ports.get(port) ports = settings.ports.get(port)
if not ports: # 意外情况 if not ports: # 意外情况
logger.log('ERROR', 'The specified request port range is incorrect') logger.log('ERROR', 'The specified request port range is incorrect')
ports = {80} ports = {80}
@@ -81,22 +81,22 @@ async def fetch(session, method, url):
""" """
timeout = aiohttp.ClientTimeout(total=None, timeout = aiohttp.ClientTimeout(total=None,
connect=None, connect=None,
sock_read=setting.sockread_timeout, sock_read=settings.sockread_timeout,
sock_connect=setting.sockconn_timeout) sock_connect=settings.sockconn_timeout)
try: try:
if method == 'HEAD': if method == 'HEAD':
async with session.head(url, async with session.head(url,
ssl=setting.verify_ssl, ssl=settings.verify_ssl,
allow_redirects=setting.allow_redirects, allow_redirects=settings.allow_redirects,
timeout=timeout, timeout=timeout,
proxy=setting.aiohttp_proxy) as resp: proxy=settings.aiohttp_proxy) as resp:
text = await resp.text() text = await resp.text()
else: else:
async with session.get(url, async with session.get(url,
ssl=setting.verify_ssl, ssl=settings.verify_ssl,
allow_redirects=setting.allow_redirects, allow_redirects=settings.allow_redirects,
timeout=timeout, timeout=timeout,
proxy=setting.aiohttp_proxy) as resp: proxy=settings.aiohttp_proxy) as resp:
try: try:
# 先尝试用utf-8解码 # 先尝试用utf-8解码
@@ -182,9 +182,9 @@ def request_callback(future, index, datas):
def get_connector(): def get_connector():
limit_open_conn = get_limit_conn() limit_open_conn = get_limit_conn()
return aiohttp.TCPConnector(ttl_dns_cache=300, return aiohttp.TCPConnector(ttl_dns_cache=300,
ssl=setting.verify_ssl, ssl=settings.verify_ssl,
limit=limit_open_conn, limit=limit_open_conn,
limit_per_host=setting.limit_per_host) limit_per_host=settings.limit_per_host)
async def async_request(urls): async def async_request(urls):
@@ -211,7 +211,7 @@ async def bulk_request(data, port):
ports = get_ports(port) ports = get_ports(port)
no_req_data = utils.get_filtered_data(data) no_req_data = utils.get_filtered_data(data)
to_req_data = gen_req_data(data, ports) to_req_data = gen_req_data(data, ports)
method = setting.request_method.upper() method = settings.request_method.upper()
logger.log('INFOR', f'Use {method} method to request') logger.log('INFOR', f'Use {method} method to request')
logger.log('INFOR', 'Async subdomains request in progress') logger.log('INFOR', 'Async subdomains request in progress')
connector = get_connector() connector = get_connector()
+5 -5
View File
@@ -2,7 +2,7 @@ import gc
import json import json
from config.log import logger from config.log import logger
from config import setting from config import settings
from common import utils from common import utils
from common.database import Database from common.database import Database
from common.ipasn import IPAsnInfo from common.ipasn import IPAsnInfo
@@ -72,7 +72,7 @@ def deal_output(output_path):
records = dict() # 用来记录所有域名解析数据 records = dict() # 用来记录所有域名解析数据
ip_asn = IPAsnInfo() ip_asn = IPAsnInfo()
ip_geo = IpGeoInfo ip_geo = IpGeoInfo
db_path = setting.data_storage_dir.joinpath('ip2region.db') db_path = settings.data_storage_dir.joinpath('ip2region.db')
ip_reg = IpRegInfo(db_path) ip_reg = IpRegInfo(db_path)
with open(output_path) as fd: with open(output_path) as fd:
for line in fd: for line in fd:
@@ -165,8 +165,8 @@ def run_resolve(domain, data):
if not subdomains: if not subdomains:
return data return data
massdns_dir = setting.third_party_dir.joinpath('massdns') massdns_dir = settings.third_party_dir.joinpath('massdns')
result_dir = setting.result_save_dir result_dir = settings.result_save_dir
temp_dir = result_dir.joinpath('temp') temp_dir = result_dir.joinpath('temp')
utils.check_dir(temp_dir) utils.check_dir(temp_dir)
massdns_path = utils.get_massdns_path(massdns_dir) massdns_path = utils.get_massdns_path(massdns_dir)
@@ -182,7 +182,7 @@ def run_resolve(domain, data):
output_path = temp_dir.joinpath(output_name) output_path = temp_dir.joinpath(output_name)
log_path = result_dir.joinpath('massdns.log') log_path = result_dir.joinpath('massdns.log')
ns_path = setting.brute_nameservers_path ns_path = settings.brute_nameservers_path
logger.log('INFOR', f'Running massdns to resolve subdomains') logger.log('INFOR', f'Running massdns to resolve subdomains')
utils.call_massdns(massdns_path, save_path, ns_path, utils.call_massdns(massdns_path, save_path, ns_path,
+4 -4
View File
@@ -1,6 +1,6 @@
import re import re
from config import setting from config import settings
from config.log import logger from config.log import logger
from common.module import Module from common.module import Module
@@ -13,8 +13,8 @@ class Search(Module):
Module.__init__(self) Module.__init__(self)
self.page_num = 0 # 要显示搜索起始条数 self.page_num = 0 # 要显示搜索起始条数
self.per_page_num = 50 # 每页显示搜索条数 self.per_page_num = 50 # 每页显示搜索条数
self.recursive_search = setting.enable_recursive_search self.recursive_search = settings.enable_recursive_search
self.recursive_times = setting.search_recursive_times self.recursive_times = settings.search_recursive_times
@staticmethod @staticmethod
def filter(domain, subdomain): def filter(domain, subdomain):
@@ -29,7 +29,7 @@ class Search(Module):
""" """
statements_list = [] statements_list = []
subdomains_temp = set(map(lambda x: x + '.' + domain, subdomains_temp = set(map(lambda x: x + '.' + domain,
setting.subdomains_common)) settings.subdomains_common))
subdomains_temp = list(subdomain.intersection(subdomains_temp)) subdomains_temp = list(subdomain.intersection(subdomains_temp))
for i in range(0, len(subdomains_temp), 2): # 同时排除2个子域 for i in range(0, len(subdomains_temp), 2): # 同时排除2个子域
statements_list.append(''.join(set(map(lambda s: ' -site:' + s, statements_list.append(''.join(set(map(lambda s: ' -site:' + s,
+16 -16
View File
@@ -19,7 +19,7 @@ from dns.resolver import Resolver
from common.domain import Domain from common.domain import Domain
from common.database import Database from common.database import Database
from config import setting from config import settings
from config.log import logger from config.log import logger
user_agents = [ user_agents = [
@@ -70,7 +70,7 @@ def get_random_header():
Get random proxy Get random proxy
""" """
header = None header = None
if setting.fake_header: if settings.fake_header:
header = gen_fake_header() header = gen_fake_header()
return header return header
@@ -80,7 +80,7 @@ def get_random_proxy():
Get random proxy Get random proxy
""" """
try: try:
return random.choice(setting.proxy_pool) return random.choice(settings.proxy_pool)
except IndexError: except IndexError:
return None return None
@@ -89,7 +89,7 @@ def get_proxy():
""" """
Get proxy Get proxy
""" """
if setting.enable_proxy: if settings.enable_proxy:
return get_random_proxy() return get_random_proxy()
return None return None
@@ -176,7 +176,7 @@ def check_path(path, name, format):
:return: 保存路径 :return: 保存路径
""" """
filename = f'{name}.{format}' filename = f'{name}.{format}'
default_path = setting.result_save_dir.joinpath(filename) default_path = settings.result_save_dir.joinpath(filename)
if isinstance(path, str): if isinstance(path, str):
path = repr(path).replace('\\', '/') # 将路径中的反斜杠替换为正斜杠 path = repr(path).replace('\\', '/') # 将路径中的反斜杠替换为正斜杠
path = path.replace('\'', '') # 去除多余的转义 path = path.replace('\'', '') # 去除多余的转义
@@ -396,9 +396,9 @@ def dns_resolver():
dns解析器 dns解析器
""" """
resolver = Resolver() resolver = Resolver()
resolver.nameservers = setting.resolver_nameservers resolver.nameservers = settings.resolver_nameservers
resolver.timeout = setting.resolver_timeout resolver.timeout = settings.resolver_timeout
resolver.lifetime = setting.resolver_lifetime resolver.lifetime = settings.resolver_lifetime
return resolver return resolver
@@ -495,7 +495,7 @@ def ip_is_public(ip_str):
def get_process_num(): def get_process_num():
process_num = setting.brute_process_num process_num = settings.brute_process_num
if isinstance(process_num, int): if isinstance(process_num, int):
return min(os.cpu_count(), process_num) return min(os.cpu_count(), process_num)
else: else:
@@ -503,7 +503,7 @@ def get_process_num():
def get_coroutine_num(): def get_coroutine_num():
coroutine_num = setting.resolve_coroutine_num coroutine_num = settings.resolve_coroutine_num
if isinstance(coroutine_num, int): if isinstance(coroutine_num, int):
return max(64, coroutine_num) return max(64, coroutine_num)
elif coroutine_num is None: elif coroutine_num is None:
@@ -601,8 +601,8 @@ def check_version(local):
api = 'https://api.github.com/repos/shmilylty/OneForAll/releases/latest' api = 'https://api.github.com/repos/shmilylty/OneForAll/releases/latest'
header = get_random_header() header = get_random_header()
proxy = get_proxy() proxy = get_proxy()
timeout = setting.request_timeout timeout = settings.request_timeout
verify = setting.request_verify verify = settings.request_verify
try: try:
resp = requests.get(url=api, headers=header, proxies=proxy, resp = requests.get(url=api, headers=header, proxies=proxy,
timeout=timeout, verify=verify) timeout=timeout, verify=verify)
@@ -633,9 +633,9 @@ def call_massdns(massdns_path, dict_path, ns_path, output_path, log_path,
quiet = '' quiet = ''
if quiet_mode: if quiet_mode:
quiet = '--quiet' quiet = '--quiet'
status_format = setting.brute_status_format status_format = settings.brute_status_format
socket_num = setting.brute_socket_num socket_num = settings.brute_socket_num
resolve_num = setting.brute_resolve_num resolve_num = settings.brute_resolve_num
cmd = f'{massdns_path} {quiet} --status-format {status_format} ' \ cmd = f'{massdns_path} {quiet} --status-format {status_format} ' \
f'--processes {process_num} --socket-count {socket_num} ' \ f'--processes {process_num} --socket-count {socket_num} ' \
f'--hashmap-size {concurrent_num} --resolvers {ns_path} ' \ f'--hashmap-size {concurrent_num} --resolvers {ns_path} ' \
@@ -648,7 +648,7 @@ def call_massdns(massdns_path, dict_path, ns_path, output_path, log_path,
def get_massdns_path(massdns_dir): def get_massdns_path(massdns_dir):
path = setting.brute_massdns_path path = settings.brute_massdns_path
if path: if path:
return path return path
system = platform.system().lower() system = platform.system().lower()
+3 -3
View File
@@ -13,7 +13,7 @@ from http.cookies import SimpleCookie
from common import utils from common import utils
from common.module import Module from common.module import Module
from config import setting from config import settings
from config.log import logger from config.log import logger
@@ -31,7 +31,7 @@ class MultiIdentify(Module):
done_queue = Manager().Queue() done_queue = Manager().Queue()
for d in data: for d in data:
task_queue.put(d) task_queue.put(d)
processes_num = min(setting.banner_process_number, os.cpu_count()) processes_num = min(settings.banner_process_number, os.cpu_count())
logger.log('INFOR', f'Creating {processes_num} processes to identify') logger.log('INFOR', f'Creating {processes_num} processes to identify')
result_data = [] result_data = []
_p = [] _p = []
@@ -54,7 +54,7 @@ class MultiIdentify(Module):
class Identify(object): class Identify(object):
def __init__(self): def __init__(self):
self.start = time.time() # 模块开始执行时间 self.start = time.time() # 模块开始执行时间
self.rule_dir = setting.data_storage_dir.joinpath('rules') self.rule_dir = settings.data_storage_dir.joinpath('rules')
self._targets = {} self._targets = {}
self.rules_num, self.RULES, self.RULE_TYPES = self.load_rules() self.rules_num, self.RULES, self.RULE_TYPES = self.load_rules()
self._cond_parser = Condition() self._cond_parser = Condition()
+6 -6
View File
@@ -4,7 +4,7 @@ import importlib
import dbexport import dbexport
from config.log import logger from config.log import logger
from config import setting from config import settings
from common import utils from common import utils
@@ -26,7 +26,7 @@ class Collect(object):
""" """
Get modules Get modules
""" """
if setting.enable_all_module: if settings.enable_all_module:
# modules = ['brute', 'certificates', 'crawl', # modules = ['brute', 'certificates', 'crawl',
# 'datasets', 'intelligence', 'search'] # 'datasets', 'intelligence', 'search']
# The crawl module has some problems # The crawl module has some problems
@@ -34,13 +34,13 @@ class Collect(object):
'dnsquery', 'intelligence', 'search'] 'dnsquery', 'intelligence', 'search']
# modules = ['certificates'] # modules = ['certificates']
for module in modules: for module in modules:
module_path = setting.module_dir.joinpath(module) module_path = settings.module_dir.joinpath(module)
for path in module_path.rglob('*.py'): for path in module_path.rglob('*.py'):
# Classes to be imported # Classes to be imported
import_module = ('modules.' + module, path.stem) import_module = ('modules.' + module, path.stem)
self.modules.append(import_module) self.modules.append(import_module)
else: else:
self.modules = setting.enable_partial_module self.modules = settings.enable_partial_module
def import_func(self): def import_func(self):
""" """
@@ -76,7 +76,7 @@ class Collect(object):
for thread in threads: for thread in threads:
# 挨个线程判断超时 最坏情况主线程阻塞时间=线程数*module_thread_timeout # 挨个线程判断超时 最坏情况主线程阻塞时间=线程数*module_thread_timeout
# 超时线程将脱离主线程 由于创建线程时已添加守护属于 所有超时线程会随着主线程结束 # 超时线程将脱离主线程 由于创建线程时已添加守护属于 所有超时线程会随着主线程结束
thread.join(setting.module_thread_timeout) thread.join(settings.module_thread_timeout)
for thread in threads: for thread in threads:
if thread.is_alive(): if thread.is_alive():
@@ -86,7 +86,7 @@ class Collect(object):
if self.export: if self.export:
if not self.path: if not self.path:
name = f'{self.domain}.{self.format}' name = f'{self.domain}.{self.format}'
self.path = setting.result_save_dir.joinpath(name) self.path = settings.result_save_dir.joinpath(name)
dbexport.export(self.domain, path=self.path, format=self.format) dbexport.export(self.domain, path=self.path, format=self.format)
end = time.time() end = time.time()
self.elapse = round(end - start, 1) self.elapse = round(end - start, 1)
+2 -2
View File
@@ -7,7 +7,7 @@ from common import utils
from common import resolve from common import resolve
from common import request from common import request
from common.module import Module from common.module import Module
from config import setting from config import settings
from config.log import logger from config.log import logger
@@ -137,7 +137,7 @@ def filter_url(domain, url, black_name):
def get_black_name(): def get_black_name():
path = setting.data_storage_dir.joinpath('common_js_library.json') path = settings.data_storage_dir.joinpath('common_js_library.json')
with open(path) as fp: with open(path) as fp:
return json.load(fp) return json.load(fp)
+2 -2
View File
@@ -1,11 +1,11 @@
import json import json
import ipaddress import ipaddress
from config import setting from config import settings
from common import utils from common import utils
from config.log import logger from config.log import logger
data_dir = setting.data_storage_dir data_dir = settings.data_storage_dir
# from https://github.com/al0ne/Vxscan/blob/master/lib/iscdn.py # from https://github.com/al0ne/Vxscan/blob/master/lib/iscdn.py
cdn_ip_cidr = utils.load_json(data_dir.joinpath('cdn_ip_cidr.json')) cdn_ip_cidr = utils.load_json(data_dir.joinpath('cdn_ip_cidr.json'))
+13 -13
View File
@@ -19,7 +19,7 @@ from common.database import Database
from modules.collect import Collect from modules.collect import Collect
from modules.finder import Finder from modules.finder import Finder
from modules import iscdn, banner from modules import iscdn, banner
from config import setting from config import settings
from config.log import logger from config.log import logger
from takeover import Takeover from takeover import Takeover
@@ -106,21 +106,21 @@ class OneForAll(object):
Configuration parameter Configuration parameter
""" """
if self.brute is None: if self.brute is None:
self.brute = bool(setting.enable_brute_module) self.brute = bool(settings.enable_brute_module)
if self.dns is None: if self.dns is None:
self.dns = bool(setting.enable_dns_resolve) self.dns = bool(settings.enable_dns_resolve)
if self.req is None: if self.req is None:
self.req = bool(setting.enable_http_request) self.req = bool(settings.enable_http_request)
if self.takeover is None: if self.takeover is None:
self.takeover = bool(setting.enable_takeover_check) self.takeover = bool(settings.enable_takeover_check)
if self.port is None: if self.port is None:
self.port = setting.http_request_port self.port = settings.http_request_port
if self.alive is None: if self.alive is None:
self.alive = bool(setting.result_export_alive) self.alive = bool(settings.result_export_alive)
if self.format is None: if self.format is None:
self.format = setting.result_save_format self.format = settings.result_save_format
if self.path is None: if self.path is None:
self.path = setting.result_save_path self.path = settings.result_save_path
def export(self, table): def export(self, table):
""" """
@@ -212,17 +212,17 @@ class OneForAll(object):
request.save_db(self.domain, self.data) request.save_db(self.domain, self.data)
# Finder module # Finder module
if setting.enable_finder_module: if settings.enable_finder_module:
finder = Finder() finder = Finder()
self.data = finder.run(self.domain, self.data, self.port) self.data = finder.run(self.domain, self.data, self.port)
# check cdn module # check cdn module
if setting.enable_cdn_check: if settings.enable_cdn_check:
self.data = iscdn.check_cdn(self.data) self.data = iscdn.check_cdn(self.data)
iscdn.save_db(self.domain, self.data) iscdn.save_db(self.domain, self.data)
# Identify banner module # Identify banner module
if setting.enable_banner_identify: if settings.enable_banner_identify:
identifier = banner.MultiIdentify() identifier = banner.MultiIdentify()
self.data = identifier.run(self.data) self.data = identifier.run(self.data)
banner.save_db(self.domain, self.data) banner.save_db(self.domain, self.data)
@@ -251,7 +251,7 @@ class OneForAll(object):
dt = datetime.now().strftime('%Y-%m-%d %H:%M:%S') dt = datetime.now().strftime('%Y-%m-%d %H:%M:%S')
print(f'[*] Starting OneForAll @ {dt}\n') print(f'[*] Starting OneForAll @ {dt}\n')
utils.check_env() utils.check_env()
if setting.enable_check_version: if settings.enable_check_version:
utils.check_version(version) utils.check_version(version)
logger.log('DEBUG', 'Python ' + utils.python_version()) logger.log('DEBUG', 'Python ' + utils.python_version())
logger.log('DEBUG', 'OneForAll ' + version) logger.log('DEBUG', 'OneForAll ' + version)
+2 -2
View File
@@ -17,14 +17,14 @@ from tablib import Dataset
from tqdm import tqdm from tqdm import tqdm
from config.log import logger from config.log import logger
from config import setting from config import settings
from common import utils from common import utils
from common.module import Module from common.module import Module
from common.domain import Domain from common.domain import Domain
def get_fingerprint(): def get_fingerprint():
path = setting.data_storage_dir.joinpath('fingerprints.json') path = settings.data_storage_dir.joinpath('fingerprints.json')
with open(path, encoding='utf-8', errors='ignore') as file: with open(path, encoding='utf-8', errors='ignore') as file:
fingerprints = json.load(file) fingerprints = json.load(file)
return fingerprints return fingerprints