mirror of
https://github.com/shmilylty/OneForAll.git
synced 2026-08-26 04:47:48 +08:00
1.重构请求,边请求边存入数据库,解决内存占用过大问题。
2.format参数改为fmt。 3.添加信息富化模块。 4.优化iscdn模块代码。 5.移除数据库new字段。
This commit is contained in:
+48
-22
@@ -57,7 +57,6 @@ class Database(object):
|
||||
f'alive int,'
|
||||
f'request int,'
|
||||
f'resolve int,'
|
||||
f'new int,'
|
||||
f'url text,'
|
||||
f'subdomain text,'
|
||||
f'port int,'
|
||||
@@ -86,6 +85,20 @@ class Database(object):
|
||||
f'elapse float,'
|
||||
f'find int)')
|
||||
|
||||
def insert_table(self, table_name, result):
|
||||
table_name = table_name.replace('.', '_')
|
||||
self.conn.query(
|
||||
f'insert into "{table_name}" '
|
||||
f'(id, alive, resolve, request, url, subdomain, port, level,'
|
||||
f'cname, ip, public, cdn, status, reason, title, banner, header,'
|
||||
f'history, response, times, ttl, cidr, asn, org, addr, isp, resolver,'
|
||||
f'module, source, elapse, find) '
|
||||
f'values (:id, :alive, :resolve, :request, :url,'
|
||||
f':subdomain, :port, :level, :cname, :ip, :public, :cdn,'
|
||||
f':status, :reason, :title, :banner, :header, :history, :response,'
|
||||
f':times, :ttl, :cidr, :asn, :org, :addr, :isp, :resolver, :module,'
|
||||
f':source, :elapse, :find)', **result)
|
||||
|
||||
def save_db(self, table_name, results, module_name=None):
|
||||
"""
|
||||
Save the results of each module in the database
|
||||
@@ -101,11 +114,11 @@ class Database(object):
|
||||
try:
|
||||
self.conn.bulk_query(
|
||||
f'insert into "{table_name}" '
|
||||
f'(id, alive, resolve, request, new, url, subdomain, port, level,'
|
||||
f'(id, alive, resolve, request, url, subdomain, port, level,'
|
||||
f'cname, ip, public, cdn, status, reason, title, banner, header,'
|
||||
f'history, response, times, ttl, cidr, asn, org, addr, isp, resolver,'
|
||||
f'module, source, elapse, find) '
|
||||
f'values (:id, :alive, :resolve, :request, :new, :url,'
|
||||
f'values (:id, :alive, :resolve, :request, :url,'
|
||||
f':subdomain, :port, :level, :cname, :ip, :public, :cdn,'
|
||||
f':status, :reason, :title, :banner, :header, :history, :response,'
|
||||
f':times, :ttl, :cidr, :asn, :org, :addr, :isp, :resolver, :module,'
|
||||
@@ -178,7 +191,7 @@ class Database(object):
|
||||
|
||||
def deduplicate_subdomain(self, table_name):
|
||||
"""
|
||||
Deduplicates of subdomains in the table
|
||||
Deduplicate subdomains in the table
|
||||
|
||||
:param str table_name: table name
|
||||
"""
|
||||
@@ -199,17 +212,6 @@ class Database(object):
|
||||
self.query(f'delete from "{table_name}" where '
|
||||
f'subdomain is null or resolve == 0')
|
||||
|
||||
def deal_table(self, deal_table_name, backup_table_name):
|
||||
"""
|
||||
Process the table when the collection task is complete
|
||||
|
||||
:param str deal_table_name: Pending table name
|
||||
:param str backup_table_name: Table name for backup
|
||||
"""
|
||||
self.copy_table(deal_table_name, backup_table_name)
|
||||
self.remove_invalid(deal_table_name)
|
||||
self.deduplicate_subdomain(deal_table_name)
|
||||
|
||||
def get_data(self, table_name):
|
||||
"""
|
||||
Get all the data in the table
|
||||
@@ -218,7 +220,7 @@ class Database(object):
|
||||
"""
|
||||
table_name = table_name.replace('.', '_')
|
||||
logger.log('TRACE', f'Get all the data from {table_name} table')
|
||||
return self.query(f'select * from "{table_name}"')
|
||||
return self.query(f'select * from {table_name}')
|
||||
|
||||
def export_data(self, table_name, alive, limit):
|
||||
"""
|
||||
@@ -229,18 +231,42 @@ class Database(object):
|
||||
:param str limit: limit value
|
||||
"""
|
||||
table_name = table_name.replace('.', '_')
|
||||
query = f'select id, new, alive, request, resolve, url, subdomain, level,' \
|
||||
f'cname, ip, public, cdn, port, status, reason, title, banner,' \
|
||||
f'cidr, asn, org, addr, isp, source from "{table_name}"'
|
||||
sql = f'select id, alive, request, resolve, url, subdomain, level,' \
|
||||
f'cname, ip, public, cdn, port, status, reason, title, banner,' \
|
||||
f'cidr, asn, org, addr, isp, source from {table_name} order by subdomain'
|
||||
if alive and limit:
|
||||
if limit in ['resolve', 'request']:
|
||||
where = f' where {limit} = 1'
|
||||
query += where
|
||||
sql += where
|
||||
elif alive:
|
||||
where = f' where alive = 1'
|
||||
query += where
|
||||
sql += where
|
||||
logger.log('TRACE', f'Get the data from {table_name} table')
|
||||
return self.query(query)
|
||||
return self.query(sql)
|
||||
|
||||
def count_alive(self, table_name):
|
||||
table_name = table_name.replace('.', '_')
|
||||
sql = f'select count() from {table_name} where alive = 1'
|
||||
return self.query(sql)
|
||||
|
||||
def get_resp_by_url(self, table_name, url):
|
||||
table_name = table_name.replace('.', '_')
|
||||
sql = f'select response from {table_name} where url = "{url}"'
|
||||
logger.log('TRACE', f'Get response data from {url}')
|
||||
return self.query(sql).scalar()
|
||||
|
||||
def get_data_by_fields(self, table_name, fields):
|
||||
table_name = table_name.replace('.', '_')
|
||||
field_str = ', '.join(fields)
|
||||
sql = f"select {field_str} from {table_name}"
|
||||
logger.log('TRACE', f'Get specified field data {fields} from {table_name} table')
|
||||
return self.query(sql)
|
||||
|
||||
def update_data_by_url(self, table_name, info, url):
|
||||
table_name = table_name.replace('.', '_')
|
||||
field_str = ', '.join(map(lambda kv: f'{kv[0]} = "{kv[1]}"', info.items()))
|
||||
sql = f'update {table_name} set {field_str} where url = "{url}"'
|
||||
return self.query(sql)
|
||||
|
||||
def close(self):
|
||||
"""
|
||||
|
||||
+1
-5
@@ -275,7 +275,6 @@ class Module(object):
|
||||
'alive': None,
|
||||
'request': None,
|
||||
'resolve': None,
|
||||
'new': None,
|
||||
'url': None,
|
||||
'subdomain': None,
|
||||
'port': None,
|
||||
@@ -315,25 +314,22 @@ class Module(object):
|
||||
ip = info.get('ip')
|
||||
times = info.get('times')
|
||||
ttl = info.get('ttl')
|
||||
public = info.get('public')
|
||||
if isinstance(cname, list):
|
||||
cname = ','.join(cname)
|
||||
ip = ','.join(ip)
|
||||
times = ','.join([str(num) for num in times])
|
||||
ttl = ','.join([str(num) for num in ttl])
|
||||
public = ','.join([str(num) for num in public])
|
||||
result = {'id': None,
|
||||
'alive': info.get('alive'),
|
||||
'request': info.get('request'),
|
||||
'resolve': info.get('resolve'),
|
||||
'new': None,
|
||||
'url': url,
|
||||
'subdomain': subdomain,
|
||||
'port': 80,
|
||||
'level': level,
|
||||
'cname': cname,
|
||||
'ip': ip,
|
||||
'public': public,
|
||||
'public': info.get('public'),
|
||||
'cdn': info.get('cdn'),
|
||||
'status': None,
|
||||
'reason': info.get('reason'),
|
||||
|
||||
+57
-54
@@ -1,7 +1,6 @@
|
||||
import json
|
||||
from threading import Thread
|
||||
from queue import Queue
|
||||
from operator import attrgetter
|
||||
|
||||
import tqdm
|
||||
import requests
|
||||
@@ -9,6 +8,7 @@ from bs4 import BeautifulSoup
|
||||
|
||||
from common import utils
|
||||
from config.log import logger
|
||||
from common.database import Database
|
||||
from config import settings
|
||||
|
||||
|
||||
@@ -127,7 +127,7 @@ def get_progress_bar(total):
|
||||
return bar
|
||||
|
||||
|
||||
def get(url, resp_list, session):
|
||||
def get_resp(url, session):
|
||||
timeout = settings.request_timeout_second
|
||||
redirect = settings.request_allow_redirect
|
||||
proxy = utils.get_proxy()
|
||||
@@ -136,13 +136,14 @@ def get(url, resp_list, session):
|
||||
except Exception as e:
|
||||
logger.log('DEBUG', e.args)
|
||||
resp = e
|
||||
resp_list.append((url, resp))
|
||||
return resp
|
||||
|
||||
|
||||
def request(urls_queue, resp_list, session):
|
||||
def request(urls_queue, resp_queue, session):
|
||||
while not urls_queue.empty():
|
||||
url = urls_queue.get()
|
||||
get(url, resp_list, session)
|
||||
index, url = urls_queue.get()
|
||||
resp = get_resp(url, session)
|
||||
resp_queue.put((index, resp))
|
||||
urls_queue.task_done()
|
||||
|
||||
|
||||
@@ -168,31 +169,6 @@ def get_session():
|
||||
return session
|
||||
|
||||
|
||||
def bulk_request(urls):
|
||||
logger.log('INFOR', 'Requesting urls in bulk')
|
||||
resp_list = list()
|
||||
urls_queue = Queue()
|
||||
for url in urls:
|
||||
urls_queue.put(url)
|
||||
total = len(urls)
|
||||
session = get_session()
|
||||
thread_count = req_thread_count()
|
||||
bar = get_progress_bar(total)
|
||||
|
||||
progress_thread = Thread(target=progress, name='ProgressThread',
|
||||
args=(bar, total, urls_queue), daemon=True)
|
||||
progress_thread.start()
|
||||
|
||||
for i in range(thread_count):
|
||||
request_thread = Thread(target=request, name=f'RequestThread-{i}',
|
||||
args=(urls_queue, resp_list, session), daemon=True)
|
||||
request_thread.start()
|
||||
|
||||
urls_queue.join()
|
||||
|
||||
return resp_list
|
||||
|
||||
|
||||
def gen_new_info(info, resp):
|
||||
if isinstance(resp, Exception):
|
||||
info['reason'] = str(resp.args)
|
||||
@@ -220,13 +196,54 @@ def gen_new_info(info, resp):
|
||||
return info
|
||||
|
||||
|
||||
def gen_new_data(data, resp_list):
|
||||
new_data = list()
|
||||
for url, resp in resp_list:
|
||||
for info in data:
|
||||
if info.get('url') == url:
|
||||
new_data.append(gen_new_info(info, resp))
|
||||
return new_data
|
||||
def save(name, total, req_data, resp_queue):
|
||||
db = Database()
|
||||
db.create_table(name)
|
||||
i = 0
|
||||
while True:
|
||||
if not resp_queue.empty():
|
||||
i += 1
|
||||
index, resp = resp_queue.get()
|
||||
old_info = req_data[index]
|
||||
new_info = gen_new_info(old_info, resp)
|
||||
db.insert_table(name, new_info)
|
||||
resp_queue.task_done()
|
||||
if i >= total:
|
||||
break
|
||||
db.close()
|
||||
|
||||
|
||||
def bulk_request(domain, req_data, ret=False):
|
||||
logger.log('INFOR', 'Requesting urls in bulk')
|
||||
resp_queue = Queue()
|
||||
urls_queue = Queue()
|
||||
total = len(req_data)
|
||||
for index, info in enumerate(req_data):
|
||||
try:
|
||||
url = info.get('url')
|
||||
except Exception:
|
||||
pass
|
||||
urls_queue.put((index, url))
|
||||
session = get_session()
|
||||
thread_count = req_thread_count()
|
||||
bar = get_progress_bar(total)
|
||||
|
||||
progress_thread = Thread(target=progress, name='ProgressThread',
|
||||
args=(bar, total, urls_queue), daemon=True)
|
||||
progress_thread.start()
|
||||
|
||||
for i in range(thread_count):
|
||||
request_thread = Thread(target=request, name=f'RequestThread-{i}',
|
||||
args=(urls_queue, resp_queue, session), daemon=True)
|
||||
request_thread.start()
|
||||
if ret:
|
||||
urls_queue.join()
|
||||
return resp_queue
|
||||
save_thread = Thread(target=save, name=f'SaveThread',
|
||||
args=(domain, total, req_data, resp_queue), daemon=True)
|
||||
save_thread.start()
|
||||
urls_queue.join()
|
||||
save_thread.join()
|
||||
|
||||
|
||||
def run_request(domain, data, port):
|
||||
@@ -242,20 +259,6 @@ def run_request(domain, data, port):
|
||||
data = utils.set_id_none(data)
|
||||
ports = get_port_seq(port)
|
||||
req_data, req_urls = gen_req_data(data, ports)
|
||||
resp_list = bulk_request(req_urls)
|
||||
new_data = gen_new_data(req_data, resp_list)
|
||||
count = utils.count_alive(new_data)
|
||||
bulk_request(domain, req_data)
|
||||
count = utils.count_alive(domain)
|
||||
logger.log('INFOR', f'Found that {domain} has {count} alive subdomains')
|
||||
sorted_data = utils.sort_by_subdomain(new_data)
|
||||
return sorted_data
|
||||
|
||||
|
||||
def save_db(name, data):
|
||||
"""
|
||||
Save request results to database
|
||||
|
||||
:param str name: table name
|
||||
:param list data: data to be saved
|
||||
"""
|
||||
logger.log('INFOR', f'Saving requested results')
|
||||
utils.save_db(name, data, 'request')
|
||||
|
||||
@@ -4,12 +4,6 @@ import json
|
||||
from config.log import logger
|
||||
from config import settings
|
||||
from common import utils
|
||||
from common.ipasn import IPAsnInfo
|
||||
from common.ipreg import IpRegData
|
||||
|
||||
|
||||
ip_asn = IPAsnInfo()
|
||||
ip_reg = IpRegData()
|
||||
|
||||
|
||||
def filter_subdomain(data):
|
||||
@@ -79,13 +73,7 @@ def gen_infos(data, qname, info, infos):
|
||||
flag = False
|
||||
cname = list()
|
||||
ips = list()
|
||||
public = list()
|
||||
ttl = list()
|
||||
cidr = list()
|
||||
asn = list()
|
||||
org = list()
|
||||
addr = list()
|
||||
isp = list()
|
||||
answers = data.get('answers')
|
||||
for answer in answers:
|
||||
if answer.get('type') == 'A':
|
||||
@@ -94,25 +82,11 @@ def gen_infos(data, qname, info, infos):
|
||||
ip = answer.get('data')
|
||||
ips.append(ip)
|
||||
ttl.append(str(answer.get('ttl')))
|
||||
public.append(str(utils.ip_is_public(ip)))
|
||||
asn_info = ip_asn.find(ip)
|
||||
cidr.append(asn_info.get('cidr'))
|
||||
asn.append(asn_info.get('asn'))
|
||||
org.append(asn_info.get('org'))
|
||||
ip_info = ip_reg.query(ip)
|
||||
addr.append(ip_info.get('addr'))
|
||||
isp.append(ip_info.get('isp'))
|
||||
info['resolve'] = 1
|
||||
info['reason'] = 'OK'
|
||||
info['cname'] = ','.join(cname)
|
||||
info['ip'] = ','.join(ips)
|
||||
info['public'] = ','.join(public)
|
||||
info['ttl'] = ','.join(ttl)
|
||||
info['cidr'] = ','.join(cidr)
|
||||
info['asn'] = ','.join(asn)
|
||||
info['org'] = ','.join(org)
|
||||
info['addr'] = ','.join(addr)
|
||||
info['isp'] = ','.join(isp)
|
||||
infos[qname] = info
|
||||
if not flag:
|
||||
info['alive'] = 0
|
||||
|
||||
+44
-17
@@ -176,16 +176,16 @@ def check_dir(dir_path):
|
||||
dir_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
|
||||
def check_path(path, name, format):
|
||||
def check_path(path, name, fmt):
|
||||
"""
|
||||
检查结果输出目录路径
|
||||
|
||||
:param path: 保存路径
|
||||
:param name: 导出名字
|
||||
:param format: 保存格式
|
||||
:param fmt: 保存格式
|
||||
:return: 保存路径
|
||||
"""
|
||||
filename = f'{name}.{format}'
|
||||
filename = f'{name}.{fmt}'
|
||||
default_path = settings.result_save_dir.joinpath(filename)
|
||||
if isinstance(path, str):
|
||||
path = repr(path).replace('\\', '/') # 将路径中的反斜杠替换为正斜杠
|
||||
@@ -204,19 +204,18 @@ def check_path(path, name, format):
|
||||
return path
|
||||
|
||||
|
||||
def check_format(format, count):
|
||||
def check_format(fmt):
|
||||
"""
|
||||
检查导出格式
|
||||
|
||||
:param format: 传入的导出格式
|
||||
:param count: 数量
|
||||
:param fmt: 传入的导出格式
|
||||
:return: 导出格式
|
||||
"""
|
||||
formats = ['csv', 'json', ]
|
||||
if format in formats:
|
||||
return format
|
||||
if fmt in formats:
|
||||
return fmt
|
||||
else:
|
||||
logger.log('ALERT', f'Does not support {format} format')
|
||||
logger.log('ALERT', f'Does not support {fmt} format')
|
||||
logger.log('ALERT', 'So use csv format by default')
|
||||
return 'csv'
|
||||
|
||||
@@ -333,8 +332,8 @@ def remove_invalid_string(string):
|
||||
return re.sub(r'[\000-\010]|[\013-\014]|[\016-\037]', r'', string)
|
||||
|
||||
|
||||
def export_all_results(path, name, format, datas):
|
||||
path = check_path(path, name, format)
|
||||
def export_all_results(path, name, fmt, datas):
|
||||
path = check_path(path, name, fmt)
|
||||
logger.log('ALERT', f'The subdomain result for all main domains: {path}')
|
||||
row_list = list()
|
||||
for row in datas:
|
||||
@@ -346,7 +345,7 @@ def export_all_results(path, name, format, datas):
|
||||
values = row.values()
|
||||
row_list.append(Record(keys, values))
|
||||
rows = RecordCollection(iter(row_list))
|
||||
content = rows.export(format)
|
||||
content = rows.export(fmt)
|
||||
save_data(path, content)
|
||||
|
||||
|
||||
@@ -366,19 +365,19 @@ def export_all_subdomains(alive, path, name, datas):
|
||||
save_data(path, data)
|
||||
|
||||
|
||||
def export_all(alive, format, path, datas):
|
||||
def export_all(alive, fmt, path, datas):
|
||||
"""
|
||||
将所有结果数据导出
|
||||
|
||||
:param bool alive: 只导出存活子域结果
|
||||
:param str format: 导出文件格式
|
||||
:param str fmt: 导出文件格式
|
||||
:param str path: 导出文件路径
|
||||
:param list datas: 待导出的结果数据
|
||||
"""
|
||||
format = check_format(format, len(datas))
|
||||
fmt = check_format(fmt)
|
||||
timestamp = get_timestring()
|
||||
name = f'all_subdomain_result_{timestamp}'
|
||||
export_all_results(path, name, format, datas)
|
||||
export_all_results(path, name, fmt, datas)
|
||||
export_all_subdomains(alive, path, name, datas)
|
||||
|
||||
|
||||
@@ -430,10 +429,18 @@ def python_version():
|
||||
return sys.version
|
||||
|
||||
|
||||
def count_alive(data):
|
||||
def calc_alive(data):
|
||||
return len(list(filter(lambda item: item.get('alive') == 1, data)))
|
||||
|
||||
|
||||
def count_alive(name):
|
||||
db = Database()
|
||||
result = db.count_alive(name)
|
||||
count = result.scalar()
|
||||
db.close()
|
||||
return count
|
||||
|
||||
|
||||
def get_subdomains(data):
|
||||
return set(map(lambda item: item.get('subdomain'), data))
|
||||
|
||||
@@ -833,3 +840,23 @@ def looks_like_ip(maybe_ip):
|
||||
return True
|
||||
except socket.error:
|
||||
return False
|
||||
|
||||
|
||||
def deal_data(domain):
|
||||
db = Database()
|
||||
db.remove_invalid(domain)
|
||||
db.deduplicate_subdomain(domain)
|
||||
db.close()
|
||||
|
||||
|
||||
def get_data(domain):
|
||||
db = Database()
|
||||
data = db.get_data(domain).as_dict()
|
||||
db.close()
|
||||
return data
|
||||
|
||||
|
||||
def clear_data(domain):
|
||||
db = Database()
|
||||
db.drop_table(domain)
|
||||
db.close()
|
||||
|
||||
Reference in New Issue
Block a user