mirror of
https://github.com/shmilylty/OneForAll.git
synced 2026-08-26 04:47:48 +08:00
重构子域请求模块
This commit is contained in:
+4
-4
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user