Files
news/tests/test_crawler.py
2026-07-18 15:51:01 +08:00

334 lines
10 KiB
Python

"""M1 抓取模块单元测试。
不依赖真实网络:用 mock 替换 Crawl4AI 的 arun。
"""
from __future__ import annotations
import asyncio
import json
from datetime import date
from pathlib import Path
from typing import Any
from unittest.mock import AsyncMock
import pytest
from crawler import engine
from crawler.config import load_crawler_config
from crawler.engine import (
crawl_url_with_retry,
extract_article_links,
)
from crawler.models import CrawlerSettings, CrawlResult, CrawlStage, SourceConfig
from crawler.storage import build_output_dir, save_result, url_hash
# --------------------------------------------------------------------------- #
# 配置加载
# --------------------------------------------------------------------------- #
def test_load_crawler_config_ok(sample_sources_yaml: Path) -> None:
cfg = load_crawler_config(sample_sources_yaml)
assert cfg.settings.concurrency == 2
assert len(cfg.sources) == 1
assert cfg.sources[0].id == "testsrc"
assert cfg.enabled_sources()[0].id == "testsrc"
def test_load_crawler_config_real_sources_yaml() -> None:
"""项目内置的 configs/sources.yaml 必须可解析(回归保护)。"""
real = Path("configs/sources.yaml")
if not real.is_file():
pytest.skip("configs/sources.yaml 未生成,跳过")
cfg = load_crawler_config(real)
# 验收标准: 至少 5 个启用源
assert len(cfg.enabled_sources()) >= 5, "启用源应至少 5 个(M1 验收标准)"
assert cfg.settings.concurrency == 3, "用户决策: 并发上限 3"
def test_load_crawler_config_missing_file(tmp_path: Path) -> None:
with pytest.raises(FileNotFoundError):
load_crawler_config(tmp_path / "nonexistent.yaml")
def test_source_id_validation() -> None:
with pytest.raises(ValueError):
SourceConfig(
id="Bad-ID", # 含大写与连字符
name="x",
homepage="https://example.com",
article_url_pattern="^.*$",
)
# --------------------------------------------------------------------------- #
# 链接抽取
# --------------------------------------------------------------------------- #
def _src(**kwargs: Any) -> SourceConfig:
base: dict[str, Any] = {
"id": "testsrc",
"name": "测试",
"homepage": "https://example.com/list",
"article_url_pattern": r"^https://example\.com/article/\d+$",
"js_render": False,
"max_articles_per_run": 10,
}
base.update(kwargs)
return SourceConfig(**base)
def test_extract_article_links_basic() -> None:
html = """
<html><body>
<a href="/article/123">文章一</a>
<a href="https://example.com/article/456">文章二</a>
<a href="https://other.com/article/789">外站</a>
<a href="/about">关于</a>
<a href="javascript:void(0)">JS</a>
<a href="/article/123">文章一(重复)</a>
</body></html>
"""
links = extract_article_links(html, "https://example.com/list", _src())
urls = [link.url for link in links]
assert urls == [
"https://example.com/article/123",
"https://example.com/article/456",
]
assert links[0].anchor_text == "文章一"
def test_extract_article_links_respects_max() -> None:
html_parts = [
f'<a href="https://example.com/article/{i}">a{i}</a>' for i in range(20)
]
html = "<html><body>" + "".join(html_parts) + "</body></html>"
src = _src(max_articles_per_run=3)
links = extract_article_links(html, "https://example.com/list", src)
assert len(links) == 3
def test_extract_article_links_strips_fragment() -> None:
html = '<a href="https://example.com/article/1#section">x</a>'
links = extract_article_links(html, "https://example.com/list", _src())
assert links[0].url == "https://example.com/article/1"
# --------------------------------------------------------------------------- #
# 重试机制
# --------------------------------------------------------------------------- #
class _FakeC4Result:
"""模拟 Crawl4AI 的返回对象。"""
def __init__(
self,
success: bool = True,
html: str = "<html>ok</html>",
markdown: str = "ok",
status_code: int = 200,
error_message: str | None = None,
) -> None:
self.success = success
self.html = html
self.markdown = markdown
self.status_code = status_code
self.error_message = error_message
@pytest.mark.asyncio
async def test_retry_succeeds_on_third_attempt() -> None:
"""前两次失败,第三次成功;返回的 attempts 应为 3。"""
fake_crawler = AsyncMock()
fake_crawler.arun = AsyncMock(
side_effect=[
_FakeC4Result(success=False, html="", error_message="boom-1"),
_FakeC4Result(success=False, html="", error_message="boom-2"),
_FakeC4Result(success=True),
]
)
settings = CrawlerSettings(
concurrency=1,
retry_max_attempts=3,
retry_min_wait_sec=0.0,
retry_max_wait_sec=0.0,
)
sem = asyncio.Semaphore(1)
src = _src()
res = await crawl_url_with_retry(
crawler=fake_crawler,
url="https://example.com/article/1",
source=src,
stage=CrawlStage.ARTICLE,
settings=settings,
semaphore=sem,
)
assert res.success is True
assert res.attempts == 3
assert fake_crawler.arun.await_count == 3
@pytest.mark.asyncio
async def test_retry_gives_up_after_max() -> None:
fake_crawler = AsyncMock()
fake_crawler.arun = AsyncMock(
return_value=_FakeC4Result(success=False, html="", error_message="nope")
)
settings = CrawlerSettings(
concurrency=1,
retry_max_attempts=2,
retry_min_wait_sec=0.0,
retry_max_wait_sec=0.0,
)
sem = asyncio.Semaphore(1)
res = await crawl_url_with_retry(
crawler=fake_crawler,
url="https://example.com/article/1",
source=_src(),
stage=CrawlStage.ARTICLE,
settings=settings,
semaphore=sem,
)
assert res.success is False
assert res.attempts == 2
assert fake_crawler.arun.await_count == 2
assert res.error == "nope"
@pytest.mark.asyncio
async def test_exception_is_swallowed_and_retried() -> None:
"""arun 抛异常应被捕获并触发重试。"""
fake_crawler = AsyncMock()
fake_crawler.arun = AsyncMock(
side_effect=[RuntimeError("net down"), _FakeC4Result(success=True)]
)
settings = CrawlerSettings(
concurrency=1, retry_max_attempts=2, retry_min_wait_sec=0.0, retry_max_wait_sec=0.0
)
sem = asyncio.Semaphore(1)
res = await crawl_url_with_retry(
crawler=fake_crawler,
url="https://example.com/article/1",
source=_src(),
stage=CrawlStage.ARTICLE,
settings=settings,
semaphore=sem,
)
assert res.success is True
assert res.attempts == 2
# --------------------------------------------------------------------------- #
# 存储
# --------------------------------------------------------------------------- #
def test_url_hash_stable() -> None:
h1 = url_hash("https://example.com/a")
h2 = url_hash("https://example.com/a")
h3 = url_hash("https://example.com/b")
assert h1 == h2
assert h1 != h3
assert len(h1) == 16
def test_build_output_dir(tmp_path: Path) -> None:
d = build_output_dir(tmp_path, "cls", date(2026, 6, 16))
assert d == tmp_path / "cls" / "20260616"
def test_save_result_writes_html_md_and_index(tmp_path: Path) -> None:
result = CrawlResult(
source_id="cls",
stage=CrawlStage.ARTICLE,
url="https://example.com/article/1",
success=True,
status_code=200,
title="标题",
html="<html>hi</html>",
markdown="# hi",
)
html_path = save_result(result, tmp_path, day=date(2026, 6, 16))
assert html_path is not None and html_path.is_file()
assert html_path.read_text(encoding="utf-8") == "<html>hi</html>"
md_path = html_path.with_suffix(".md")
assert md_path.is_file()
assert md_path.read_text(encoding="utf-8") == "# hi"
index = html_path.parent / "index.jsonl"
assert index.is_file()
line = index.read_text(encoding="utf-8").strip()
obj = json.loads(line)
assert obj["url"] == "https://example.com/article/1"
assert obj["success"] is True
assert obj["html_file"] == html_path.name
assert "html" not in obj # 大字段不应进入元数据
def test_save_result_failure_only_appends_index(tmp_path: Path) -> None:
result = CrawlResult(
source_id="cls",
stage=CrawlStage.ARTICLE,
url="https://example.com/article/2",
success=False,
error="timeout",
)
html_path = save_result(result, tmp_path, day=date(2026, 6, 16))
assert html_path is None
index = tmp_path / "cls" / "20260616" / "index.jsonl"
assert index.is_file()
obj = json.loads(index.read_text(encoding="utf-8").strip())
assert obj["success"] is False
# --------------------------------------------------------------------------- #
# Markdown 兼容
# --------------------------------------------------------------------------- #
def test_markdown_text_handles_str() -> None:
assert engine._markdown_text("plain") == "plain"
def test_markdown_text_handles_object_raw_markdown() -> None:
class _Obj:
raw_markdown = "from raw"
assert engine._markdown_text(_Obj()) == "from raw"
def test_markdown_text_handles_none() -> None:
assert engine._markdown_text(None) == ""
# --------------------------------------------------------------------------- #
# 集成测试(默认跳过,需要真实浏览器与网络)
# --------------------------------------------------------------------------- #
@pytest.mark.integration
@pytest.mark.asyncio
async def test_real_homepage_crawl_smoke() -> None:
"""真实抓取 example.com 烟测,验证端到端可运行。
运行: uv run pytest -m integration
"""
from crawler import crawl_all
from crawler.models import CrawlerConfig
cfg = CrawlerConfig(
settings=CrawlerSettings(concurrency=1, retry_max_attempts=1, output_root="data/raw_test"),
sources=[
SourceConfig(
id="example",
name="example",
homepage="https://example.com/",
article_url_pattern=r"^https://www\.iana\.org/.*$",
js_render=False,
page_timeout_ms=15000,
max_articles_per_run=1,
)
],
)
results = await crawl_all(cfg, save=False)
assert any(r.success for r in results), "example.com 烟测应至少一个成功"