重构子域请求模块

This commit is contained in:
Jing Ling
2020-08-26 03:39:58 +08:00
parent a4d6872921
commit dc1ca6cffc
13 changed files with 202 additions and 439 deletions
+4 -4
View File
@@ -22,9 +22,9 @@ class Module(object):
self.cookie = None
self.header = dict()
self.proxy = None
self.delay = settings.request_delay # 请求睡眠时延
self.timeout = settings.request_timeout # 请求超时时间
self.verify = settings.request_verify # 请求SSL验证
self.delay = 1 # 请求睡眠时延
self.timeout = settings.request_timeout_second # 请求超时时间
self.verify = settings.request_ssl_verify # 请求SSL验证
self.domain = str() # 当前进行子域名收集的主域
self.subdomains = set() # 存放发现的子域
self.infos = dict() # 存放子域有关信息
@@ -199,7 +199,7 @@ class Module(object):
:param str module: module name
:return: proxy
"""
if not settings.enable_proxy:
if not settings.enable_request_proxy:
logger.log('TRACE', f'All modules do not use proxy')
return self.proxy
if settings.proxy_all_module:
+132 -187
View File
@@ -1,10 +1,10 @@
import json
import asyncio
import functools
from threading import Thread
from queue import Queue
import aiohttp
import tqdm
from aiohttp import ClientSession
import requests
from bs4 import BeautifulSoup
from common import utils
@@ -12,17 +12,17 @@ from config.log import logger
from config import settings
def get_limit_conn():
count = settings.limit_open_conn
def req_thread_count():
count = settings.request_thread_count
if isinstance(count, int):
count = max(16, count)
else:
count = utils.get_coroutine_count()
logger.log('DEBUG', f'Request coroutine {count}')
count = utils.get_request_count()
logger.log('DEBUG', f'Number of request threads {count}')
return count
def get_ports(port):
def get_port_seq(port):
logger.log('DEBUG', 'Getting port range')
ports = set()
if isinstance(port, (set, list, tuple)):
@@ -40,81 +40,33 @@ def get_ports(port):
return set(ports)
def gen_req_data(data, ports):
def gen_req_url(domain, port):
if str(port).endswith('443'):
url = f'https://{domain}:{port}'
if port == 443:
url = f'https://{domain}'
return url
url = f'http://{domain}:{port}'
if port == 80:
url = f'http://{domain}'
return url
def gen_req_urls(data, ports):
logger.log('INFOR', 'Generating request urls')
new_data = []
for data in data:
resolve = data.get('resolve')
urls = set()
for info in data:
resolve = info.get('resolve')
# 解析不成功的子域不进行http请求探测
if resolve != 1:
continue
subdomain = data.get('subdomain')
subdomain = info.get('subdomain')
for port in ports:
if str(port).endswith('443'):
url = f'https://{subdomain}:{port}'
if port == 443:
url = f'https://{subdomain}'
data['id'] = None
data['url'] = url
data['port'] = port
new_data.append(data)
data = dict(data) # 需要生成一个新的字典对象
else:
url = f'http://{subdomain}:{port}'
if port == 80:
url = f'http://{subdomain}'
data['id'] = None
data['url'] = url
data['port'] = port
new_data.append(data)
data = dict(data) # 需要生成一个新的字典对象
return new_data
urls.add(gen_req_url(subdomain, port))
return urls
async def fetch(session, method, url):
"""
请求
:param session: session对象
:param method: 请求方法
:param str url: url地址
:return: 响应对象和响应文本
"""
timeout = aiohttp.ClientTimeout(total=None,
connect=None,
sock_read=settings.sockread_timeout,
sock_connect=settings.sockconn_timeout)
try:
if method == 'HEAD':
async with session.head(url,
ssl=settings.verify_ssl,
allow_redirects=settings.allow_redirects,
timeout=timeout,
proxy=settings.aiohttp_proxy) as resp:
text = await resp.text()
else:
async with session.get(url,
ssl=settings.verify_ssl,
allow_redirects=settings.allow_redirects,
timeout=timeout,
proxy=settings.aiohttp_proxy) as resp:
try:
# 先尝试用utf-8解码
text = await resp.text(encoding='utf-8', errors='strict')
except UnicodeError:
try:
# 再尝试用gb18030解码
text = await resp.text(encoding='gb18030', errors='strict')
except UnicodeError:
# 最后尝试自动解码
text = await resp.text(encoding=None, errors='ignore')
return resp, text
except Exception as e:
return e, None
def get_title(markup):
def get_html_title(markup):
"""
获取标题
@@ -161,107 +113,107 @@ def get_jump_urls(history):
return urls
def request_callback(future, index, datas):
resp, text = future.result()
if isinstance(resp, BaseException):
exception = resp
logger.log('TRACE', exception.args)
name = utils.get_classname(exception)
datas[index]['reason'] = name + ' ' + str(exception)
datas[index]['request'] = 0
datas[index]['alive'] = 0
else:
datas[index]['reason'] = resp.reason
datas[index]['status'] = resp.status
datas[index]['request'] = 1
if resp.status == 400 or resp.status >= 500:
datas[index]['alive'] = 0
else:
datas[index]['alive'] = 1
headers = resp.headers
if settings.enable_banner_identify:
datas[index]['banner'] = utils.get_sample_banner(headers)
datas[index]['header'] = json.dumps(dict(headers))
history = resp.history
datas[index]['history'] = json.dumps(get_jump_urls(history))
if isinstance(text, str):
title = get_title(text).strip()
datas[index]['title'] = utils.remove_invalid_string(title)
datas[index]['response'] = utils.remove_invalid_string(text)
def get_connector():
count = get_limit_conn()
return aiohttp.TCPConnector(ttl_dns_cache=300,
ssl=settings.verify_ssl,
limit=count,
limit_per_host=settings.limit_per_host)
async def async_request(urls):
results = list()
connector = get_connector()
headers = utils.get_random_header()
session = ClientSession(connector=connector, headers=headers)
tasks = []
for i, url in enumerate(urls):
task = asyncio.ensure_future(fetch(session, 'GET', url))
tasks.append(task)
if tasks:
futures = asyncio.as_completed(tasks)
for future in tqdm.tqdm(futures,
total=len(tasks),
desc='Request Progress',
ncols=80):
result = await future
results.append(result)
await session.close()
return results
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 = settings.request_method.upper()
logger.log('INFOR', f'Use {method} method to request')
logger.log('INFOR', 'Async subdomains request in progress')
connector = get_connector()
headers = utils.get_random_header()
session = ClientSession(connector=connector, headers=headers)
tasks = []
for num, data in enumerate(to_req_data):
url = data.get('url')
task = asyncio.ensure_future(fetch(session, method, url))
task.add_done_callback(functools.partial(request_callback,
index=num,
datas=to_req_data))
tasks.append(task)
if tasks:
futures = asyncio.as_completed(tasks)
for future in tqdm.tqdm(futures,
total=len(tasks),
desc='Request Progress',
ncols=80):
await future
await session.close()
return to_req_data + no_req_data
def set_loop_policy():
def get(url, resp_queue, session):
timeout = settings.request_timeout_second
redirect = settings.request_allow_redirect
proxy = utils.get_proxy()
try:
import uvloop
except ImportError:
pass
resp = session.get(url, timeout=timeout, allow_redirects=redirect, proxies=proxy)
except Exception as e:
logger.log('DEBUG', e.args)
return
resp_queue.put(resp)
def request(urls_queue, resp_queue, session):
while not urls_queue.empty():
url = urls_queue.get()
get(url, resp_queue, session)
urls_queue.task_done()
def progress(urls, resp_queue):
bar = tqdm.tqdm()
bar.total = len(urls)
bar.desc = 'Request Progress'
bar.ncols = 80
while True:
done = resp_queue.qsize()
bar.n = done
bar.update()
if done == bar.total:
break
def get_session():
header = utils.gen_fake_header()
verify = settings.request_ssl_verify
redirect_limit = settings.request_redirect_limit
session = requests.Session()
session.headers = header
session.verify = verify
session.max_redirects = redirect_limit
return session
def bulk_request(urls):
logger.log('INFOR', 'Requesting urls in bulk')
resp_list = list()
urls_queue = Queue()
resp_queue = Queue()
for url in urls:
urls_queue.put(url)
session = get_session()
thread_count = req_thread_count()
progress_thread = Thread(target=progress, args=(urls, resp_queue))
progress_thread.start()
for _ in range(thread_count):
request_thread = Thread(target=request, args=(urls_queue, resp_queue, session))
request_thread.start()
urls_queue.join()
while not resp_queue.empty():
resp = resp_queue.get()
resp_list.append(resp)
return resp_list
def gen_new_info(info, resp):
port = resp.raw._pool.port
info['port'] = port
info['url'] = resp.url
info['reason'] = resp.reason
code = resp.status_code
info['status'] = code
info['request'] = 1
if code == 400 or code >= 500:
info['alive'] = 0
else:
asyncio.set_event_loop_policy(uvloop.EventLoopPolicy())
info['alive'] = 1
headers = resp.headers
if settings.enable_banner_identify:
info['banner'] = utils.get_sample_banner(headers)
info['header'] = json.dumps(dict(headers))
history = resp.history
info['history'] = json.dumps(get_jump_urls(history))
text = utils.decode_resp_text(resp)
title = get_html_title(text).strip()
info['title'] = utils.remove_invalid_string(title)
info['response'] = utils.remove_invalid_string(text)
return info
def set_loop():
set_loop_policy()
loop = asyncio.get_event_loop()
asyncio.set_event_loop(loop)
return loop
def gen_new_data(data, resp_list):
new_data = list()
for resp in resp_list:
subdomain = resp.raw._pool.host
for info in data:
if info.get('subdomain') == subdomain:
new_data.append(gen_new_info(info, resp))
return new_data
def run_request(domain, data, port):
@@ -274,25 +226,18 @@ def run_request(domain, data, port):
:return list: result
"""
logger.log('INFOR', f'Start requesting subdomains of {domain}')
loop = set_loop()
data = utils.set_id_none(data)
request_coroutine = bulk_request(data, port)
data = loop.run_until_complete(request_coroutine)
loop.run_until_complete(asyncio.sleep(0.25))
ports = get_port_seq(port)
filtered_data = utils.get_filtered_data(data)
req_urls = gen_req_urls(data, ports)
resp_list = bulk_request(req_urls)
new_data = gen_new_data(data, resp_list)
data = new_data + filtered_data
count = utils.count_alive(data)
logger.log('INFOR', f'Found that {domain} has {count} alive subdomains')
return data
def urls_request(urls):
logger.log('INFOR', 'Start urls request module')
loop = set_loop()
request_coroutine = async_request(urls)
data = loop.run_until_complete(request_coroutine)
loop.run_until_complete(asyncio.sleep(0.25))
return data
def save_db(name, data):
"""
Save request results to database
+19 -39
View File
@@ -10,7 +10,6 @@ import subprocess
from ipaddress import IPv4Address, ip_address
from stat import S_IXUSR
import psutil
import tenacity
import requests
from pathlib import Path
@@ -49,10 +48,10 @@ def gen_fake_header():
"""
Generate fake request headers
"""
headers = settings.headers
headers = settings.request_default_headers
if not isinstance(headers, dict):
headers = dict()
if settings.random_user_agent:
if settings.enable_random_ua:
ua = random.choice(user_agents)
headers['User-Agent'] = ua
headers['Accept-Encoding'] = 'gzip, deflate'
@@ -74,7 +73,7 @@ def get_random_proxy():
Get random proxy
"""
try:
return random.choice(settings.proxy_pool)
return random.choice(settings.request_proxy_pool)
except IndexError:
return None
@@ -83,7 +82,7 @@ def get_proxy():
"""
Get proxy
"""
if settings.enable_proxy:
if settings.enable_request_proxy:
return get_random_proxy()
return None
@@ -502,27 +501,8 @@ def get_process_num():
return 1
def get_coroutine_count():
"""
根据内存大小获取并发数量
"""
mem = psutil.virtual_memory()
total_mem = mem.total
g_size = 1024 * 1024 * 1024
if total_mem < 1 * g_size:
return 16
elif total_mem < 2 * g_size:
return 32
elif total_mem < 4 * g_size:
return 64
elif total_mem < 8 * g_size:
return 128
elif total_mem < 16 * g_size:
return 256
elif total_mem < 32 * g_size:
return 512
else:
return 1024
def get_request_count():
return os.cpu_count() * 10
def uniq_dict_list(dict_list):
@@ -562,7 +542,6 @@ def check_net():
def check_pre():
logger.log('INFOR', 'Checking dependent environment')
system = platform.system()
implementation = platform.python_implementation()
version = platform.python_version()
if implementation != 'CPython':
@@ -571,16 +550,6 @@ def check_pre():
if version < '3.6':
logger.log('FATAL', 'OneForAll requires Python 3.6 or higher')
exit(1)
if system == 'Windows' and implementation == 'CPython' and version < '3.8':
logger.log('FATAL', 'OneForAll requires Python 3.8 '
'or higher when running on Windows')
exit(1)
if system in {"Linux", "Darwin"}:
try:
import uvloop
except ImportError:
logger.log('ALERT', f'Please install the uvloop library manually '
f'to accelerate subdomain requests')
def check_env():
@@ -599,8 +568,8 @@ def check_version(local):
api = 'https://api.github.com/repos/shmilylty/OneForAll/releases/latest'
header = get_random_header()
proxy = get_proxy()
timeout = settings.request_timeout
verify = settings.request_verify
timeout = settings.request_timeout_second
verify = settings.request_ssl_verify
try:
resp = requests.get(url=api, headers=header, proxies=proxy,
timeout=timeout, verify=verify)
@@ -747,3 +716,14 @@ def get_url_resp(url):
logger.log('DEBUG', e.args)
return None
return resp
def decode_resp_text(resp):
try:
text = resp.text(encoding='utf-8', errors='strict') # 先尝试用utf-8严格解码
except UnicodeError:
try:
text = resp.text(encoding='gb18030', errors='strict') # 再尝试用gb18030严格解码
except UnicodeError:
text = resp.text(encoding=None, errors='ignore') # 最后尝试自动解码
return text