实现配置文件插拔式设计

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