1.重构请求,边请求边存入数据库,解决内存占用过大问题。

2.format参数改为fmt。
3.添加信息富化模块。
4.优化iscdn模块代码。
5.移除数据库new字段。
This commit is contained in:
Jing Ling
2020-11-09 09:01:06 +08:00
parent 2428017d29
commit 3fbbd1dc25
17 changed files with 473 additions and 424 deletions
+48 -22
View File
@@ -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
View File
@@ -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
View File
@@ -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')
-26
View File
@@ -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
View File
@@ -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()