Test refactoring, part 2. (#6928)

This commit is contained in:
Andrey Rakhmatullin 2025-06-30 02:10:38 +05:00 committed by GitHub
parent 2c1c10e923
commit db0be1771c
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
12 changed files with 501 additions and 434 deletions

View File

@ -104,10 +104,8 @@ class TestSelector:
Selector(TextResponse(url="http://example.com", body=b""), text="") Selector(TextResponse(url="http://example.com", body=b""), text="")
@pytest.mark.skipif(not PARSEL_18_PLUS, reason="parsel < 1.8 doesn't support jmespath")
class TestJMESPath: class TestJMESPath:
@pytest.mark.skipif(
not PARSEL_18_PLUS, reason="parsel < 1.8 doesn't support jmespath"
)
def test_json_has_html(self) -> None: def test_json_has_html(self) -> None:
"""Sometimes the information is returned in a json wrapper""" """Sometimes the information is returned in a json wrapper"""
@ -145,9 +143,6 @@ class TestJMESPath:
assert resp.jmespath("html").css("div > b").getall() == ["<b>f</b>"] assert resp.jmespath("html").css("div > b").getall() == ["<b>f</b>"]
assert resp.jmespath("content").jmespath("name.age").get() == "18" assert resp.jmespath("content").jmespath("name.age").get() == "18"
@pytest.mark.skipif(
not PARSEL_18_PLUS, reason="parsel < 1.8 doesn't support jmespath"
)
def test_html_has_json(self) -> None: def test_html_has_json(self) -> None:
body = """ body = """
<div> <div>
@ -193,9 +188,6 @@ class TestJMESPath:
] ]
assert resp.xpath("//div/content").jmespath("total").get() == "4" assert resp.xpath("//div/content").jmespath("total").get() == "4"
@pytest.mark.skipif(
not PARSEL_18_PLUS, reason="parsel < 1.8 doesn't support jmespath"
)
def test_jmestpath_with_re(self) -> None: def test_jmestpath_with_re(self) -> None:
body = """ body = """
<div> <div>
@ -248,13 +240,14 @@ class TestJMESPath:
r"(\d+)" r"(\d+)"
) == ["18", "32", "22", "25"] ) == ["18", "32", "22", "25"]
@pytest.mark.skipif(PARSEL_18_PLUS, reason="parsel >= 1.8 supports jmespath")
def test_jmespath_not_available(self) -> None: @pytest.mark.skipif(PARSEL_18_PLUS, reason="parsel >= 1.8 supports jmespath")
body = """ def test_jmespath_not_available() -> None:
{ body = """
"website": {"name": "Example"} {
} "website": {"name": "Example"}
""" }
resp = TextResponse(url="http://example.com", body=body, encoding="utf-8") """
with pytest.raises(AttributeError): resp = TextResponse(url="http://example.com", body=body, encoding="utf-8")
resp.jmespath("website.name").get() with pytest.raises(AttributeError):
resp.jmespath("website.name").get()

View File

@ -1,3 +1,5 @@
from __future__ import annotations
import gzip import gzip
import warnings import warnings
from io import BytesIO from io import BytesIO
@ -531,7 +533,7 @@ class TestSitemapSpider(TestSpider):
g.close() g.close()
GZBODY = f.getvalue() GZBODY = f.getvalue()
def assertSitemapBody(self, response, body): def assertSitemapBody(self, response: Response, body: bytes | None) -> None:
crawler = get_crawler() crawler = get_crawler()
spider = self.spider_class.from_crawler(crawler, "example.com") spider = self.spider_class.from_crawler(crawler, "example.com")
assert spider._get_sitemap_body(response) == body assert spider._get_sitemap_body(response) == body

View File

@ -1,5 +1,8 @@
from __future__ import annotations
import warnings import warnings
from asyncio import sleep from asyncio import sleep
from typing import Any
import pytest import pytest
from testfixtures import LogCapture from testfixtures import LogCapture
@ -19,7 +22,9 @@ ITEM_B = {"id": "b"}
class TestMain(TestCase): class TestMain(TestCase):
async def _test_spider(self, spider, expected_items=None): async def _test_spider(
self, spider: type[Spider], expected_items: list[Any] | None = None
) -> None:
actual_items = [] actual_items = []
expected_items = [] if expected_items is None else expected_items expected_items = [] if expected_items is None else expected_items
@ -29,6 +34,7 @@ class TestMain(TestCase):
crawler = get_crawler(spider) crawler = get_crawler(spider)
crawler.signals.connect(track_item, signals.item_scraped) crawler.signals.connect(track_item, signals.item_scraped)
await maybe_deferred_to_future(crawler.crawl()) await maybe_deferred_to_future(crawler.crawl())
assert crawler.stats
assert crawler.stats.get_value("finish_reason") == "finished" assert crawler.stats.get_value("finish_reason") == "finished"
assert actual_items == expected_items assert actual_items == expected_items

View File

@ -15,10 +15,7 @@ from scrapy.exceptions import _InvalidOutput
from scrapy.http import Request, Response from scrapy.http import Request, Response
from scrapy.spiders import Spider from scrapy.spiders import Spider
from scrapy.utils.asyncgen import collect_asyncgen from scrapy.utils.asyncgen import collect_asyncgen
from scrapy.utils.defer import ( from scrapy.utils.defer import deferred_f_from_coro_f, maybe_deferred_to_future
deferred_f_from_coro_f,
maybe_deferred_to_future,
)
from scrapy.utils.test import get_crawler from scrapy.utils.test import get_crawler
@ -118,7 +115,9 @@ class TestBaseAsyncSpiderMiddleware(TestSpiderMiddleware):
RESULT_COUNT = 3 # to simplify checks, let everything return 3 objects RESULT_COUNT = 3 # to simplify checks, let everything return 3 objects
@staticmethod @staticmethod
def _construct_mw_setting(*mw_classes, start_index: int | None = None): def _construct_mw_setting(
*mw_classes: type[Any], start_index: int | None = None
) -> dict[type[Any], int]:
if start_index is None: if start_index is None:
start_index = 10 start_index = 10
return {i: c for c, i in enumerate(mw_classes, start=start_index)} return {i: c for c, i in enumerate(mw_classes, start=start_index)}
@ -128,7 +127,9 @@ class TestBaseAsyncSpiderMiddleware(TestSpiderMiddleware):
yield {"foo": 2} yield {"foo": 2}
yield {"foo": 3} yield {"foo": 3}
async def _get_middleware_result(self, *mw_classes, start_index: int | None = None): async def _get_middleware_result(
self, *mw_classes: type[Any], start_index: int | None = None
) -> Any:
setting = self._construct_mw_setting(*mw_classes, start_index=start_index) setting = self._construct_mw_setting(*mw_classes, start_index=start_index)
self.crawler = get_crawler( self.crawler = get_crawler(
Spider, {"SPIDER_MIDDLEWARES_BASE": {}, "SPIDER_MIDDLEWARES": setting} Spider, {"SPIDER_MIDDLEWARES_BASE": {}, "SPIDER_MIDDLEWARES": setting}
@ -140,8 +141,11 @@ class TestBaseAsyncSpiderMiddleware(TestSpiderMiddleware):
) )
async def _test_simple_base( async def _test_simple_base(
self, *mw_classes, downgrade: bool = False, start_index: int | None = None self,
): *mw_classes: type[Any],
downgrade: bool = False,
start_index: int | None = None,
) -> None:
with LogCapture() as log: with LogCapture() as log:
result = await self._get_middleware_result( result = await self._get_middleware_result(
*mw_classes, start_index=start_index *mw_classes, start_index=start_index
@ -156,8 +160,11 @@ class TestBaseAsyncSpiderMiddleware(TestSpiderMiddleware):
) )
async def _test_asyncgen_base( async def _test_asyncgen_base(
self, *mw_classes, downgrade: bool = False, start_index: int | None = None self,
): *mw_classes: type[Any],
downgrade: bool = False,
start_index: int | None = None,
) -> None:
with LogCapture() as log: with LogCapture() as log:
result = await self._get_middleware_result( result = await self._get_middleware_result(
*mw_classes, start_index=start_index *mw_classes, start_index=start_index
@ -335,7 +342,9 @@ class TestProcessStartSimple(TestBaseAsyncSpiderMiddleware):
ITEM_TYPE = (Request, dict) ITEM_TYPE = (Request, dict)
MW_SIMPLE = ProcessStartSimpleMiddleware MW_SIMPLE = ProcessStartSimpleMiddleware
async def _get_processed_start(self, *mw_classes): async def _get_processed_start(
self, *mw_classes: type[Any]
) -> AsyncIterator[Any] | None:
class TestSpider(Spider): class TestSpider(Spider):
name = "test" name = "test"
@ -384,61 +393,65 @@ class UniversalMiddlewareBothAsync:
class TestUniversalMiddlewareManager: class TestUniversalMiddlewareManager:
def setup_method(self): @pytest.fixture
self.mwman = SpiderMiddlewareManager() def mwman(self) -> SpiderMiddlewareManager:
return SpiderMiddlewareManager()
def test_simple_mw(self): def test_simple_mw(self, mwman: SpiderMiddlewareManager) -> None:
mw = ProcessSpiderOutputSimpleMiddleware() mw = ProcessSpiderOutputSimpleMiddleware()
self.mwman._add_middleware(mw) mwman._add_middleware(mw)
assert ( assert (
self.mwman.methods["process_spider_output"][0] == mw.process_spider_output # pylint: disable=comparison-with-callable mwman.methods["process_spider_output"][0] == mw.process_spider_output # pylint: disable=comparison-with-callable
) )
def test_async_mw(self): def test_async_mw(self, mwman: SpiderMiddlewareManager) -> None:
mw = ProcessSpiderOutputAsyncGenMiddleware() mw = ProcessSpiderOutputAsyncGenMiddleware()
self.mwman._add_middleware(mw) mwman._add_middleware(mw)
assert ( assert (
self.mwman.methods["process_spider_output"][0] == mw.process_spider_output # pylint: disable=comparison-with-callable mwman.methods["process_spider_output"][0] == mw.process_spider_output # pylint: disable=comparison-with-callable
) )
def test_universal_mw(self): def test_universal_mw(self, mwman: SpiderMiddlewareManager) -> None:
mw = ProcessSpiderOutputUniversalMiddleware() mw = ProcessSpiderOutputUniversalMiddleware()
self.mwman._add_middleware(mw) mwman._add_middleware(mw)
assert self.mwman.methods["process_spider_output"][0] == ( assert mwman.methods["process_spider_output"][0] == (
mw.process_spider_output, mw.process_spider_output,
mw.process_spider_output_async, mw.process_spider_output_async,
) )
def test_universal_mw_no_sync(self): def test_universal_mw_no_sync(
with LogCapture() as log: self, mwman: SpiderMiddlewareManager, caplog: pytest.LogCaptureFixture
self.mwman._add_middleware(UniversalMiddlewareNoSync()) ) -> None:
mwman._add_middleware(UniversalMiddlewareNoSync())
assert ( assert (
"UniversalMiddlewareNoSync has process_spider_output_async" "UniversalMiddlewareNoSync has process_spider_output_async"
" without process_spider_output" in str(log) " without process_spider_output" in caplog.text
) )
assert self.mwman.methods["process_spider_output"][0] is None assert mwman.methods["process_spider_output"][0] is None
def test_universal_mw_both_sync(self): def test_universal_mw_both_sync(
self, mwman: SpiderMiddlewareManager, caplog: pytest.LogCaptureFixture
) -> None:
mw = UniversalMiddlewareBothSync() mw = UniversalMiddlewareBothSync()
with LogCapture() as log: mwman._add_middleware(mw)
self.mwman._add_middleware(mw)
assert ( assert (
"UniversalMiddlewareBothSync.process_spider_output_async " "UniversalMiddlewareBothSync.process_spider_output_async "
"is not an async generator function" in str(log) "is not an async generator function" in caplog.text
) )
assert ( assert (
self.mwman.methods["process_spider_output"][0] == mw.process_spider_output # pylint: disable=comparison-with-callable mwman.methods["process_spider_output"][0] == mw.process_spider_output # pylint: disable=comparison-with-callable
) )
def test_universal_mw_both_async(self): def test_universal_mw_both_async(
with LogCapture() as log: self, mwman: SpiderMiddlewareManager, caplog: pytest.LogCaptureFixture
self.mwman._add_middleware(UniversalMiddlewareBothAsync()) ) -> None:
mwman._add_middleware(UniversalMiddlewareBothAsync())
assert ( assert (
"UniversalMiddlewareBothAsync.process_spider_output " "UniversalMiddlewareBothAsync.process_spider_output "
"is an async generator function while process_spider_output_async exists" "is an async generator function while process_spider_output_async exists"
in str(log) in caplog.text
) )
assert self.mwman.methods["process_spider_output"][0] is None assert mwman.methods["process_spider_output"][0] is None
class TestBuiltinMiddlewareSimple(TestBaseAsyncSpiderMiddleware): class TestBuiltinMiddlewareSimple(TestBaseAsyncSpiderMiddleware):
@ -447,7 +460,9 @@ class TestBuiltinMiddlewareSimple(TestBaseAsyncSpiderMiddleware):
MW_ASYNCGEN = ProcessSpiderOutputAsyncGenMiddleware MW_ASYNCGEN = ProcessSpiderOutputAsyncGenMiddleware
MW_UNIVERSAL = ProcessSpiderOutputUniversalMiddleware MW_UNIVERSAL = ProcessSpiderOutputUniversalMiddleware
async def _get_middleware_result(self, *mw_classes, start_index: int | None = None): async def _get_middleware_result(
self, *mw_classes: type[Any], start_index: int | None = None
) -> Any:
setting = self._construct_mw_setting(*mw_classes, start_index=start_index) setting = self._construct_mw_setting(*mw_classes, start_index=start_index)
self.crawler = get_crawler(Spider, {"SPIDER_MIDDLEWARES": setting}) self.crawler = get_crawler(Spider, {"SPIDER_MIDDLEWARES": setting})
self.spider = self.crawler._create_spider("foo") self.spider = self.crawler._create_spider("foo")
@ -534,7 +549,7 @@ class TestProcessSpiderException(TestBaseAsyncSpiderMiddleware):
def _scrape_func(self, *args, **kwargs): def _scrape_func(self, *args, **kwargs):
1 / 0 1 / 0
async def _test_asyncgen_nodowngrade(self, *mw_classes): async def _test_asyncgen_nodowngrade(self, *mw_classes: type[Any]) -> None:
with pytest.raises( with pytest.raises(
_InvalidOutput, match="Async iterable returned from .+ cannot be downgraded" _InvalidOutput, match="Async iterable returned from .+ cannot be downgraded"
): ):

View File

@ -18,7 +18,7 @@ def crawler() -> Crawler:
return get_crawler(Spider) return get_crawler(Spider)
def test_trivial(crawler): def test_trivial(crawler: Crawler) -> None:
class TrivialSpiderMiddleware(BaseSpiderMiddleware): class TrivialSpiderMiddleware(BaseSpiderMiddleware):
pass pass
@ -28,15 +28,13 @@ def test_trivial(crawler):
test_req = Request("data:,") test_req = Request("data:,")
spider_output = [test_req, {"foo": "bar"}] spider_output = [test_req, {"foo": "bar"}]
for processed in [ for processed in [
list( list(mw.process_spider_output(Response("data:,"), spider_output, None)), # type: ignore[arg-type]
mw.process_spider_output(Response("data:,"), spider_output, crawler.spider) list(mw.process_start_requests(spider_output, None)), # type: ignore[arg-type]
),
list(mw.process_start_requests(spider_output, crawler.spider)),
]: ]:
assert processed == [test_req, {"foo": "bar"}] assert processed == [test_req, {"foo": "bar"}]
def test_processed_request(crawler): def test_processed_request(crawler: Crawler) -> None:
class ProcessReqSpiderMiddleware(BaseSpiderMiddleware): class ProcessReqSpiderMiddleware(BaseSpiderMiddleware):
def get_processed_request( def get_processed_request(
self, request: Request, response: Response | None self, request: Request, response: Response | None
@ -53,10 +51,8 @@ def test_processed_request(crawler):
test_req3 = Request("data:3,") test_req3 = Request("data:3,")
spider_output = [test_req1, {"foo": "bar"}, test_req2, test_req3] spider_output = [test_req1, {"foo": "bar"}, test_req2, test_req3]
for processed in [ for processed in [
list( list(mw.process_spider_output(Response("data:,"), spider_output, None)), # type: ignore[arg-type]
mw.process_spider_output(Response("data:,"), spider_output, crawler.spider) list(mw.process_start_requests(spider_output, None)), # type: ignore[arg-type]
),
list(mw.process_start_requests(spider_output, crawler.spider)),
]: ]:
assert len(processed) == 3 assert len(processed) == 3
assert isinstance(processed[0], Request) assert isinstance(processed[0], Request)
@ -66,7 +62,7 @@ def test_processed_request(crawler):
assert processed[2].url == "data:30," assert processed[2].url == "data:30,"
def test_processed_item(crawler): def test_processed_item(crawler: Crawler) -> None:
class ProcessItemSpiderMiddleware(BaseSpiderMiddleware): class ProcessItemSpiderMiddleware(BaseSpiderMiddleware):
def get_processed_item(self, item: Any, response: Response | None) -> Any: def get_processed_item(self, item: Any, response: Response | None) -> Any:
if item["foo"] == 2: if item["foo"] == 2:
@ -79,15 +75,13 @@ def test_processed_item(crawler):
test_req = Request("data:,") test_req = Request("data:,")
spider_output = [{"foo": 1}, {"foo": 2}, test_req, {"foo": 3}] spider_output = [{"foo": 1}, {"foo": 2}, test_req, {"foo": 3}]
for processed in [ for processed in [
list( list(mw.process_spider_output(Response("data:,"), spider_output, None)), # type: ignore[arg-type]
mw.process_spider_output(Response("data:,"), spider_output, crawler.spider) list(mw.process_start_requests(spider_output, None)), # type: ignore[arg-type]
),
list(mw.process_start_requests(spider_output, crawler.spider)),
]: ]:
assert processed == [{"foo": 1}, test_req, {"foo": 30}] assert processed == [{"foo": 1}, test_req, {"foo": 30}]
def test_processed_both(crawler): def test_processed_both(crawler: Crawler) -> None:
class ProcessBothSpiderMiddleware(BaseSpiderMiddleware): class ProcessBothSpiderMiddleware(BaseSpiderMiddleware):
def get_processed_request( def get_processed_request(
self, request: Request, response: Response | None self, request: Request, response: Response | None
@ -118,10 +112,8 @@ def test_processed_both(crawler):
test_req3, test_req3,
] ]
for processed in [ for processed in [
list( list(mw.process_spider_output(Response("data:,"), spider_output, None)), # type: ignore[arg-type]
mw.process_spider_output(Response("data:,"), spider_output, crawler.spider) list(mw.process_start_requests(spider_output, None)), # type: ignore[arg-type]
),
list(mw.process_start_requests(spider_output, crawler.spider)),
]: ]:
assert len(processed) == 4 assert len(processed) == 4
assert isinstance(processed[0], Request) assert isinstance(processed[0], Request)

View File

@ -1,38 +1,65 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import pytest
from scrapy.http import Request, Response from scrapy.http import Request, Response
from scrapy.spidermiddlewares.depth import DepthMiddleware from scrapy.spidermiddlewares.depth import DepthMiddleware
from scrapy.spiders import Spider from scrapy.spiders import Spider
from scrapy.utils.test import get_crawler from scrapy.utils.test import get_crawler
if TYPE_CHECKING:
from collections.abc import Generator
class TestDepthMiddleware: from scrapy.crawler import Crawler
def setup_method(self): from scrapy.statscollectors import StatsCollector
crawler = get_crawler(Spider, {"DEPTH_LIMIT": 1, "DEPTH_STATS_VERBOSE": True})
self.spider = crawler._create_spider("scrapytest.org")
self.stats = crawler.stats
self.stats.open_spider(self.spider)
self.mw = DepthMiddleware.from_crawler(crawler) @pytest.fixture
def crawler() -> Crawler:
return get_crawler(Spider, {"DEPTH_LIMIT": 1, "DEPTH_STATS_VERBOSE": True})
def test_process_spider_output(self):
req = Request("http://scrapytest.org")
resp = Response("http://scrapytest.org")
resp.request = req
result = [Request("http://scrapytest.org")]
out = list(self.mw.process_spider_output(resp, result, self.spider)) @pytest.fixture
assert out == result def spider(crawler: Crawler) -> Spider:
crawler.spider = crawler._create_spider("scrapytest.org")
return crawler.spider
rdc = self.stats.get_value("request_depth_count/1", spider=self.spider)
assert rdc == 1
req.meta["depth"] = 1 @pytest.fixture
def stats(crawler: Crawler, spider: Spider) -> Generator[StatsCollector]:
assert crawler.stats is not None
crawler.stats.open_spider(spider)
out2 = list(self.mw.process_spider_output(resp, result, self.spider)) yield crawler.stats
assert not out2
rdm = self.stats.get_value("request_depth_max", spider=self.spider) crawler.stats.close_spider(spider, "")
assert rdm == 1
def teardown_method(self):
self.stats.close_spider(self.spider, "") @pytest.fixture
def mw(crawler: Crawler) -> DepthMiddleware:
return DepthMiddleware.from_crawler(crawler)
def test_process_spider_output(
mw: DepthMiddleware, stats: StatsCollector, spider: Spider
) -> None:
req = Request("http://scrapytest.org")
resp = Response("http://scrapytest.org")
resp.request = req
result = [Request("http://scrapytest.org")]
out = list(mw.process_spider_output(resp, result, spider))
assert out == result
rdc = stats.get_value("request_depth_count/1", spider=spider)
assert rdc == 1
req.meta["depth"] = 1
out2 = list(mw.process_spider_output(resp, result, spider))
assert not out2
rdm = stats.get_value("request_depth_max", spider=spider)
assert rdm == 1

View File

@ -1,3 +1,5 @@
from __future__ import annotations
import logging import logging
import pytest import pytest
@ -49,111 +51,144 @@ class _HttpErrorSpider(MockServerSpider):
return failure return failure
def _responses(request, status_codes): req = Request("http://scrapytest.org")
responses = []
for code in status_codes:
response = Response(request.url, status=code) def _response(request: Request, status_code: int) -> Response:
response.request = request return Response(request.url, status=status_code, request=request)
responses.append(response)
return responses
@pytest.fixture
def spider() -> Spider:
crawler = get_crawler(Spider)
return Spider.from_crawler(crawler, name="foo")
@pytest.fixture
def res200() -> Response:
return _response(req, 200)
@pytest.fixture
def res402() -> Response:
return _response(req, 402)
@pytest.fixture
def res404() -> Response:
return _response(req, 404)
class TestHttpErrorMiddleware: class TestHttpErrorMiddleware:
def setup_method(self): @pytest.fixture
crawler = get_crawler(Spider) def mw(self) -> HttpErrorMiddleware:
self.spider = Spider.from_crawler(crawler, name="foo") return HttpErrorMiddleware(Settings({}))
self.mw = HttpErrorMiddleware(Settings({}))
self.req = Request("http://scrapytest.org")
self.res200, self.res404 = _responses(self.req, [200, 404])
def test_process_spider_input(self): def test_process_spider_input(
assert self.mw.process_spider_input(self.res200, self.spider) is None self,
mw: HttpErrorMiddleware,
spider: Spider,
res200: Response,
res404: Response,
) -> None:
mw.process_spider_input(res200, spider)
with pytest.raises(HttpError): with pytest.raises(HttpError):
self.mw.process_spider_input(self.res404, self.spider) mw.process_spider_input(res404, spider)
def test_process_spider_exception(self): def test_process_spider_exception(
assert ( self, mw: HttpErrorMiddleware, spider: Spider, res404: Response
self.mw.process_spider_exception( ) -> None:
self.res404, HttpError(self.res404), self.spider assert mw.process_spider_exception(res404, HttpError(res404), spider) == []
) assert mw.process_spider_exception(res404, Exception(), spider) is None
== []
)
assert (
self.mw.process_spider_exception(self.res404, Exception(), self.spider)
is None
)
def test_handle_httpstatus_list(self): def test_handle_httpstatus_list(
res = self.res404.copy() self, mw: HttpErrorMiddleware, spider: Spider, res404: Response
res.request = Request( ) -> None:
request = Request(
"http://scrapytest.org", meta={"handle_httpstatus_list": [404]} "http://scrapytest.org", meta={"handle_httpstatus_list": [404]}
) )
assert self.mw.process_spider_input(res, self.spider) is None res = _response(request, 404)
mw.process_spider_input(res, spider)
self.spider.handle_httpstatus_list = [404] spider.handle_httpstatus_list = [404] # type: ignore[attr-defined]
assert self.mw.process_spider_input(self.res404, self.spider) is None mw.process_spider_input(res404, spider)
class TestHttpErrorMiddlewareSettings: class TestHttpErrorMiddlewareSettings:
"""Similar test, but with settings""" """Similar test, but with settings"""
def setup_method(self): @pytest.fixture
self.spider = Spider("foo") def mw(self) -> HttpErrorMiddleware:
self.mw = HttpErrorMiddleware(Settings({"HTTPERROR_ALLOWED_CODES": (402,)})) return HttpErrorMiddleware(Settings({"HTTPERROR_ALLOWED_CODES": (402,)}))
self.req = Request("http://scrapytest.org")
self.res200, self.res404, self.res402 = _responses(self.req, [200, 404, 402])
def test_process_spider_input(self): def test_process_spider_input(
assert self.mw.process_spider_input(self.res200, self.spider) is None self,
mw: HttpErrorMiddleware,
spider: Spider,
res200: Response,
res402: Response,
res404: Response,
) -> None:
mw.process_spider_input(res200, spider)
with pytest.raises(HttpError): with pytest.raises(HttpError):
self.mw.process_spider_input(self.res404, self.spider) mw.process_spider_input(res404, spider)
assert self.mw.process_spider_input(self.res402, self.spider) is None mw.process_spider_input(res402, spider)
def test_meta_overrides_settings(self): def test_meta_overrides_settings(
self, mw: HttpErrorMiddleware, spider: Spider
) -> None:
request = Request( request = Request(
"http://scrapytest.org", meta={"handle_httpstatus_list": [404]} "http://scrapytest.org", meta={"handle_httpstatus_list": [404]}
) )
res404 = self.res404.copy() res404 = _response(request, 404)
res404.request = request res402 = _response(request, 402)
res402 = self.res402.copy()
res402.request = request
assert self.mw.process_spider_input(res404, self.spider) is None mw.process_spider_input(res404, spider)
with pytest.raises(HttpError): with pytest.raises(HttpError):
self.mw.process_spider_input(res402, self.spider) mw.process_spider_input(res402, spider)
def test_spider_override_settings(self): def test_spider_override_settings(
self.spider.handle_httpstatus_list = [404] self,
assert self.mw.process_spider_input(self.res404, self.spider) is None mw: HttpErrorMiddleware,
spider: Spider,
res402: Response,
res404: Response,
) -> None:
spider.handle_httpstatus_list = [404] # type: ignore[attr-defined]
mw.process_spider_input(res404, spider)
with pytest.raises(HttpError): with pytest.raises(HttpError):
self.mw.process_spider_input(self.res402, self.spider) mw.process_spider_input(res402, spider)
class TestHttpErrorMiddlewareHandleAll: class TestHttpErrorMiddlewareHandleAll:
def setup_method(self): @pytest.fixture
self.spider = Spider("foo") def mw(self) -> HttpErrorMiddleware:
self.mw = HttpErrorMiddleware(Settings({"HTTPERROR_ALLOW_ALL": True})) return HttpErrorMiddleware(Settings({"HTTPERROR_ALLOW_ALL": True}))
self.req = Request("http://scrapytest.org")
self.res200, self.res404, self.res402 = _responses(self.req, [200, 404, 402])
def test_process_spider_input(self): def test_process_spider_input(
assert self.mw.process_spider_input(self.res200, self.spider) is None self,
assert self.mw.process_spider_input(self.res404, self.spider) is None mw: HttpErrorMiddleware,
spider: Spider,
res200: Response,
res404: Response,
) -> None:
mw.process_spider_input(res200, spider)
mw.process_spider_input(res404, spider)
def test_meta_overrides_settings(self): def test_meta_overrides_settings(
self, mw: HttpErrorMiddleware, spider: Spider
) -> None:
request = Request( request = Request(
"http://scrapytest.org", meta={"handle_httpstatus_list": [404]} "http://scrapytest.org", meta={"handle_httpstatus_list": [404]}
) )
res404 = self.res404.copy() res404 = _response(request, 404)
res404.request = request res402 = _response(request, 402)
res402 = self.res402.copy()
res402.request = request
assert self.mw.process_spider_input(res404, self.spider) is None mw.process_spider_input(res404, spider)
with pytest.raises(HttpError): with pytest.raises(HttpError):
self.mw.process_spider_input(res402, self.spider) mw.process_spider_input(res402, spider)
def test_httperror_allow_all_false(self): def test_httperror_allow_all_false(self, spider: Spider) -> None:
crawler = get_crawler(_HttpErrorSpider) crawler = get_crawler(_HttpErrorSpider)
mw = HttpErrorMiddleware.from_crawler(crawler) mw = HttpErrorMiddleware.from_crawler(crawler)
request_httpstatus_false = Request( request_httpstatus_false = Request(
@ -162,14 +197,12 @@ class TestHttpErrorMiddlewareHandleAll:
request_httpstatus_true = Request( request_httpstatus_true = Request(
"http://scrapytest.org", meta={"handle_httpstatus_all": True} "http://scrapytest.org", meta={"handle_httpstatus_all": True}
) )
res404 = self.res404.copy() res404 = _response(request_httpstatus_false, 404)
res404.request = request_httpstatus_false res402 = _response(request_httpstatus_true, 402)
res402 = self.res402.copy()
res402.request = request_httpstatus_true
with pytest.raises(HttpError): with pytest.raises(HttpError):
mw.process_spider_input(res404, self.spider) mw.process_spider_input(res404, spider)
assert mw.process_spider_input(res402, self.spider) is None mw.process_spider_input(res402, spider)
class TestHttpErrorMiddlewareIntegrational(TestCase): class TestHttpErrorMiddlewareIntegrational(TestCase):

View File

@ -1,7 +1,7 @@
from __future__ import annotations from __future__ import annotations
import warnings import warnings
from typing import Any from typing import Any, cast
from urllib.parse import urlparse from urllib.parse import urlparse
import pytest import pytest
@ -34,6 +34,11 @@ from scrapy.spidermiddlewares.referer import (
from scrapy.spiders import Spider from scrapy.spiders import Spider
@pytest.fixture
def spider() -> Spider:
return Spider("foo")
class TestRefererMiddleware: class TestRefererMiddleware:
req_meta: dict[str, Any] = {} req_meta: dict[str, Any] = {}
resp_headers: dict[str, str] = {} resp_headers: dict[str, str] = {}
@ -42,22 +47,22 @@ class TestRefererMiddleware:
("http://scrapytest.org", "http://scrapytest.org/", b"http://scrapytest.org"), ("http://scrapytest.org", "http://scrapytest.org/", b"http://scrapytest.org"),
] ]
def setup_method(self): @pytest.fixture
self.spider = Spider("foo") def mw(self) -> RefererMiddleware:
settings = Settings(self.settings) settings = Settings(self.settings)
self.mw = RefererMiddleware(settings) return RefererMiddleware(settings)
def get_request(self, target): def get_request(self, target: str) -> Request:
return Request(target, meta=self.req_meta) return Request(target, meta=self.req_meta)
def get_response(self, origin): def get_response(self, origin: str) -> Response:
return Response(origin, headers=self.resp_headers) return Response(origin, headers=self.resp_headers)
def test(self): def test(self, mw: RefererMiddleware, spider: Spider) -> None:
for origin, target, referrer in self.scenarii: for origin, target, referrer in self.scenarii:
response = self.get_response(origin) response = self.get_response(origin)
request = self.get_request(target) request = self.get_request(target)
out = list(self.mw.process_spider_output(response, [request], self.spider)) out = list(mw.process_spider_output(response, [request], spider))
assert out[0].headers.get("Referer") == referrer assert out[0].headers.get("Referer") == referrer
@ -1002,13 +1007,22 @@ class TestReferrerOnRedirect(TestRefererMiddleware):
), ),
] ]
def setup_method(self): @pytest.fixture
self.spider = Spider("foo") def referrermw(self) -> RefererMiddleware:
settings = Settings(self.settings) settings = Settings(self.settings)
self.referrermw = RefererMiddleware(settings) return RefererMiddleware(settings)
self.redirectmw = RedirectMiddleware(settings)
def test(self): @pytest.fixture
def redirectmw(self) -> RedirectMiddleware:
settings = Settings(self.settings)
return RedirectMiddleware(settings)
def test( # type: ignore[override]
self,
referrermw: RefererMiddleware,
redirectmw: RedirectMiddleware,
spider: Spider,
) -> None:
for ( for (
parent, parent,
target, target,
@ -1019,19 +1033,17 @@ class TestReferrerOnRedirect(TestRefererMiddleware):
response = self.get_response(parent) response = self.get_response(parent)
request = self.get_request(target) request = self.get_request(target)
out = list( out = list(referrermw.process_spider_output(response, [request], spider))
self.referrermw.process_spider_output(response, [request], self.spider)
)
assert out[0].headers.get("Referer") == init_referrer assert out[0].headers.get("Referer") == init_referrer
for status, url in redirections: for status, url in redirections:
response = Response( response = Response(
request.url, headers={"Location": url}, status=status request.url, headers={"Location": url}, status=status
) )
request = self.redirectmw.process_response( request = cast(
request, response, self.spider Request, redirectmw.process_response(request, response, spider)
) )
self.referrermw.request_scheduled(request, self.spider) referrermw.request_scheduled(request, spider)
assert isinstance(request, Request) assert isinstance(request, Request)
assert request.headers.get("Referer") == final_referrer assert request.headers.get("Referer") == final_referrer

View File

@ -1,39 +1,64 @@
from testfixtures import LogCapture from __future__ import annotations
from logging import INFO
from typing import TYPE_CHECKING
import pytest
from scrapy.http import Request, Response from scrapy.http import Request, Response
from scrapy.spidermiddlewares.urllength import UrlLengthMiddleware from scrapy.spidermiddlewares.urllength import UrlLengthMiddleware
from scrapy.spiders import Spider from scrapy.spiders import Spider
from scrapy.utils.test import get_crawler from scrapy.utils.test import get_crawler
if TYPE_CHECKING:
from scrapy.crawler import Crawler
from scrapy.statscollectors import StatsCollector
class TestUrlLengthMiddleware:
def setup_method(self):
self.maxlength = 25
crawler = get_crawler(Spider, {"URLLENGTH_LIMIT": self.maxlength})
self.spider = crawler._create_spider("foo")
self.stats = crawler.stats
self.mw = UrlLengthMiddleware.from_crawler(crawler)
self.response = Response("http://scrapytest.org") maxlength = 25
self.short_url_req = Request("http://scrapytest.org/") response = Response("http://scrapytest.org")
self.long_url_req = Request("http://scrapytest.org/this_is_a_long_url") short_url_req = Request("http://scrapytest.org/")
self.reqs = [self.short_url_req, self.long_url_req] long_url_req = Request("http://scrapytest.org/this_is_a_long_url")
reqs: list[Request] = [short_url_req, long_url_req]
def process_spider_output(self):
return list(
self.mw.process_spider_output(self.response, self.reqs, self.spider)
)
def test_middleware_works(self): @pytest.fixture
assert self.process_spider_output() == [self.short_url_req] def crawler() -> Crawler:
return get_crawler(Spider, {"URLLENGTH_LIMIT": maxlength})
def test_logging(self):
with LogCapture() as log:
self.process_spider_output()
ric = self.stats.get_value( @pytest.fixture
"urllength/request_ignored_count", spider=self.spider def spider(crawler: Crawler) -> Spider:
) return crawler._create_spider("foo")
assert ric == 1
assert f"Ignoring link (url length > {self.maxlength})" in str(log)
@pytest.fixture
def stats(crawler: Crawler) -> StatsCollector:
assert crawler.stats is not None
return crawler.stats
@pytest.fixture
def mw(crawler: Crawler) -> UrlLengthMiddleware:
return UrlLengthMiddleware.from_crawler(crawler)
def process_spider_output(mw: UrlLengthMiddleware, spider: Spider) -> list[Request]:
return list(mw.process_spider_output(response, reqs, spider))
def test_middleware_works(mw: UrlLengthMiddleware, spider: Spider) -> None:
assert process_spider_output(mw, spider) == [short_url_req]
def test_logging(
stats: StatsCollector,
mw: UrlLengthMiddleware,
spider: Spider,
caplog: pytest.LogCaptureFixture,
) -> None:
with caplog.at_level(INFO):
process_spider_output(mw, spider)
ric = stats.get_value("urllength/request_ignored_count", spider=spider)
assert ric == 1
assert f"Ignoring link (url length > {maxlength})" in caplog.text

View File

@ -1,6 +1,4 @@
import shutil
from datetime import datetime, timezone from datetime import datetime, timezone
from tempfile import mkdtemp
import pytest import pytest
@ -10,37 +8,36 @@ from scrapy.spiders import Spider
from scrapy.utils.test import get_crawler from scrapy.utils.test import get_crawler
class TestSpiderState: def test_store_load(tmp_path):
def test_store_load(self): jobdir = str(tmp_path)
jobdir = mkdtemp()
try:
spider = Spider(name="default")
dt = datetime.now(tz=timezone.utc)
ss = SpiderState(jobdir) spider = Spider(name="default")
ss.spider_opened(spider) dt = datetime.now(tz=timezone.utc)
spider.state["one"] = 1
spider.state["dt"] = dt
ss.spider_closed(spider)
spider2 = Spider(name="default") ss = SpiderState(jobdir)
ss2 = SpiderState(jobdir) ss.spider_opened(spider)
ss2.spider_opened(spider2) spider.state["one"] = 1
assert spider.state == {"one": 1, "dt": dt} spider.state["dt"] = dt
ss2.spider_closed(spider2) ss.spider_closed(spider)
finally:
shutil.rmtree(jobdir)
def test_state_attribute(self): spider2 = Spider(name="default")
# state attribute must be present if jobdir is not set, to provide a ss2 = SpiderState(jobdir)
# consistent interface ss2.spider_opened(spider2)
spider = Spider(name="default") assert spider.state == {"one": 1, "dt": dt}
ss = SpiderState() ss2.spider_closed(spider2)
ss.spider_opened(spider)
assert spider.state == {}
ss.spider_closed(spider)
def test_not_configured(self):
crawler = get_crawler(Spider) def test_state_attribute():
with pytest.raises(NotConfigured): # state attribute must be present if jobdir is not set, to provide a
SpiderState.from_crawler(crawler) # consistent interface
spider = Spider(name="default")
ss = SpiderState()
ss.spider_opened(spider)
assert spider.state == {}
ss.spider_closed(spider)
def test_not_configured():
crawler = get_crawler(Spider)
with pytest.raises(NotConfigured):
SpiderState.from_crawler(crawler)

View File

@ -2,7 +2,10 @@
Queues that handle requests Queues that handle requests
""" """
from pathlib import Path from __future__ import annotations
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING
import pytest import pytest
import queuelib import queuelib
@ -19,86 +22,65 @@ from scrapy.squeues import (
) )
from scrapy.utils.test import get_crawler from scrapy.utils.test import get_crawler
if TYPE_CHECKING:
class TestBaseQueue: from scrapy.crawler import Crawler
def setup_method(self):
self.crawler = get_crawler(Spider)
class RequestQueueTestMixin: HAVE_PEEK = hasattr(queuelib.queue.FifoMemoryQueue, "peek")
def queue(self, base_path: Path):
@pytest.fixture
def crawler() -> Crawler:
return get_crawler(Spider)
class TestRequestQueueBase(ABC):
@property
@abstractmethod
def is_fifo(self) -> bool:
raise NotImplementedError raise NotImplementedError
def test_one_element_with_peek(self, tmp_path): @pytest.mark.parametrize("test_peek", [True, False])
if not hasattr(queuelib.queue.FifoMemoryQueue, "peek"): def test_one_element(self, q: queuelib.queue.BaseQueue, test_peek: bool):
if test_peek and not HAVE_PEEK:
pytest.skip("The queuelib queues do not define peek") pytest.skip("The queuelib queues do not define peek")
q = self.queue(tmp_path) if not test_peek and HAVE_PEEK:
pytest.skip("The queuelib queues define peek")
assert len(q) == 0 assert len(q) == 0
assert q.peek() is None if test_peek:
assert q.peek() is None
assert q.pop() is None assert q.pop() is None
req = Request("http://www.example.com") req = Request("http://www.example.com")
q.push(req) q.push(req)
assert len(q) == 1 assert len(q) == 1
assert q.peek().url == req.url if test_peek:
assert q.pop().url == req.url result = q.peek()
assert result is not None
assert result.url == req.url
else:
with pytest.raises(
NotImplementedError,
match="The underlying queue class does not implement 'peek'",
):
q.peek()
result = q.pop()
assert result is not None
assert result.url == req.url
assert len(q) == 0 assert len(q) == 0
assert q.peek() is None if test_peek:
assert q.peek() is None
assert q.pop() is None assert q.pop() is None
q.close() q.close()
def test_one_element_without_peek(self, tmp_path): @pytest.mark.parametrize("test_peek", [True, False])
if hasattr(queuelib.queue.FifoMemoryQueue, "peek"): def test_order(self, q: queuelib.queue.BaseQueue, test_peek: bool):
pytest.skip("The queuelib queues define peek") if test_peek and not HAVE_PEEK:
q = self.queue(tmp_path)
assert len(q) == 0
assert q.pop() is None
req = Request("http://www.example.com")
q.push(req)
assert len(q) == 1
with pytest.raises(
NotImplementedError,
match="The underlying queue class does not implement 'peek'",
):
q.peek()
assert q.pop().url == req.url
assert len(q) == 0
assert q.pop() is None
q.close()
class FifoQueueMixin(RequestQueueTestMixin):
def test_fifo_with_peek(self, tmp_path):
if not hasattr(queuelib.queue.FifoMemoryQueue, "peek"):
pytest.skip("The queuelib queues do not define peek") pytest.skip("The queuelib queues do not define peek")
q = self.queue(tmp_path) if not test_peek and HAVE_PEEK:
assert len(q) == 0
assert q.peek() is None
assert q.pop() is None
req1 = Request("http://www.example.com/1")
req2 = Request("http://www.example.com/2")
req3 = Request("http://www.example.com/3")
q.push(req1)
q.push(req2)
q.push(req3)
assert len(q) == 3
assert q.peek().url == req1.url
assert q.pop().url == req1.url
assert len(q) == 2
assert q.peek().url == req2.url
assert q.pop().url == req2.url
assert len(q) == 1
assert q.peek().url == req3.url
assert q.pop().url == req3.url
assert len(q) == 0
assert q.peek() is None
assert q.pop() is None
q.close()
def test_fifo_without_peek(self, tmp_path):
if hasattr(queuelib.queue.FifoMemoryQueue, "peek"):
pytest.skip("The queuelib queues define peek") pytest.skip("The queuelib queues define peek")
q = self.queue(tmp_path)
assert len(q) == 0 assert len(q) == 0
if test_peek:
assert q.peek() is None
assert q.pop() is None assert q.pop() is None
req1 = Request("http://www.example.com/1") req1 = Request("http://www.example.com/1")
req2 = Request("http://www.example.com/2") req2 = Request("http://www.example.com/2")
@ -106,111 +88,80 @@ class FifoQueueMixin(RequestQueueTestMixin):
q.push(req1) q.push(req1)
q.push(req2) q.push(req2)
q.push(req3) q.push(req3)
with pytest.raises( if not test_peek:
NotImplementedError, with pytest.raises(
match="The underlying queue class does not implement 'peek'", NotImplementedError,
): match="The underlying queue class does not implement 'peek'",
q.peek() ):
assert len(q) == 3 q.peek()
assert q.pop().url == req1.url reqs = [req1, req2, req3] if self.is_fifo else [req3, req2, req1]
assert len(q) == 2 for i, req in enumerate(reqs):
assert q.pop().url == req2.url assert len(q) == 3 - i
assert len(q) == 1 if test_peek:
assert q.pop().url == req3.url result = q.peek()
assert result is not None
assert result.url == req.url
result = q.pop()
assert result is not None
assert result.url == req.url
assert len(q) == 0 assert len(q) == 0
if test_peek:
assert q.peek() is None
assert q.pop() is None assert q.pop() is None
q.close() q.close()
class LifoQueueMixin(RequestQueueTestMixin): class TestPickleFifoDiskQueueRequest(TestRequestQueueBase):
def test_lifo_with_peek(self, tmp_path): is_fifo = True
if not hasattr(queuelib.queue.FifoMemoryQueue, "peek"):
pytest.skip("The queuelib queues do not define peek")
q = self.queue(tmp_path)
assert len(q) == 0
assert q.peek() is None
assert q.pop() is None
req1 = Request("http://www.example.com/1")
req2 = Request("http://www.example.com/2")
req3 = Request("http://www.example.com/3")
q.push(req1)
q.push(req2)
q.push(req3)
assert len(q) == 3
assert q.peek().url == req3.url
assert q.pop().url == req3.url
assert len(q) == 2
assert q.peek().url == req2.url
assert q.pop().url == req2.url
assert len(q) == 1
assert q.peek().url == req1.url
assert q.pop().url == req1.url
assert len(q) == 0
assert q.peek() is None
assert q.pop() is None
q.close()
def test_lifo_without_peek(self, tmp_path): @pytest.fixture
if hasattr(queuelib.queue.FifoMemoryQueue, "peek"): def q(self, crawler, tmp_path):
pytest.skip("The queuelib queues define peek")
q = self.queue(tmp_path)
assert len(q) == 0
assert q.pop() is None
req1 = Request("http://www.example.com/1")
req2 = Request("http://www.example.com/2")
req3 = Request("http://www.example.com/3")
q.push(req1)
q.push(req2)
q.push(req3)
with pytest.raises(
NotImplementedError,
match="The underlying queue class does not implement 'peek'",
):
q.peek()
assert len(q) == 3
assert q.pop().url == req3.url
assert len(q) == 2
assert q.pop().url == req2.url
assert len(q) == 1
assert q.pop().url == req1.url
assert len(q) == 0
assert q.pop() is None
q.close()
class TestPickleFifoDiskQueueRequest(FifoQueueMixin, TestBaseQueue):
def queue(self, base_path):
return PickleFifoDiskQueue.from_crawler( return PickleFifoDiskQueue.from_crawler(
crawler=self.crawler, key=str(base_path / "pickle" / "fifo") crawler=crawler, key=str(tmp_path / "pickle" / "fifo")
) )
class TestPickleLifoDiskQueueRequest(LifoQueueMixin, TestBaseQueue): class TestPickleLifoDiskQueueRequest(TestRequestQueueBase):
def queue(self, base_path): is_fifo = False
@pytest.fixture
def q(self, crawler, tmp_path):
return PickleLifoDiskQueue.from_crawler( return PickleLifoDiskQueue.from_crawler(
crawler=self.crawler, key=str(base_path / "pickle" / "lifo") crawler=crawler, key=str(tmp_path / "pickle" / "lifo")
) )
class TestMarshalFifoDiskQueueRequest(FifoQueueMixin, TestBaseQueue): class TestMarshalFifoDiskQueueRequest(TestRequestQueueBase):
def queue(self, base_path): is_fifo = True
@pytest.fixture
def q(self, crawler, tmp_path):
return MarshalFifoDiskQueue.from_crawler( return MarshalFifoDiskQueue.from_crawler(
crawler=self.crawler, key=str(base_path / "marshal" / "fifo") crawler=crawler, key=str(tmp_path / "marshal" / "fifo")
) )
class TestMarshalLifoDiskQueueRequest(LifoQueueMixin, TestBaseQueue): class TestMarshalLifoDiskQueueRequest(TestRequestQueueBase):
def queue(self, base_path): is_fifo = False
@pytest.fixture
def q(self, crawler, tmp_path):
return MarshalLifoDiskQueue.from_crawler( return MarshalLifoDiskQueue.from_crawler(
crawler=self.crawler, key=str(base_path / "marshal" / "lifo") crawler=crawler, key=str(tmp_path / "marshal" / "lifo")
) )
class TestFifoMemoryQueueRequest(FifoQueueMixin, TestBaseQueue): class TestFifoMemoryQueueRequest(TestRequestQueueBase):
def queue(self, base_path): is_fifo = True
return FifoMemoryQueue.from_crawler(crawler=self.crawler)
@pytest.fixture
def q(self, crawler):
return FifoMemoryQueue.from_crawler(crawler=crawler)
class TestLifoMemoryQueueRequest(LifoQueueMixin, TestBaseQueue): class TestLifoMemoryQueueRequest(TestRequestQueueBase):
def queue(self, base_path): is_fifo = False
return LifoMemoryQueue.from_crawler(crawler=self.crawler)
@pytest.fixture
def q(self, crawler):
return LifoMemoryQueue.from_crawler(crawler=crawler)

View File

@ -1,28 +1,44 @@
from __future__ import annotations
from datetime import datetime from datetime import datetime
from typing import TYPE_CHECKING
from unittest import mock from unittest import mock
import pytest
from scrapy.extensions.corestats import CoreStats from scrapy.extensions.corestats import CoreStats
from scrapy.spiders import Spider from scrapy.spiders import Spider
from scrapy.statscollectors import DummyStatsCollector, StatsCollector from scrapy.statscollectors import DummyStatsCollector, StatsCollector
from scrapy.utils.test import get_crawler from scrapy.utils.test import get_crawler
if TYPE_CHECKING:
from scrapy.crawler import Crawler
@pytest.fixture
def crawler() -> Crawler:
return get_crawler(Spider)
@pytest.fixture
def spider(crawler: Crawler) -> Spider:
return crawler._create_spider("foo")
class TestCoreStatsExtension: class TestCoreStatsExtension:
def setup_method(self):
self.crawler = get_crawler(Spider)
self.spider = self.crawler._create_spider("foo")
@mock.patch("scrapy.extensions.corestats.datetime") @mock.patch("scrapy.extensions.corestats.datetime")
def test_core_stats_default_stats_collector(self, mock_datetime): def test_core_stats_default_stats_collector(
self, mock_datetime: mock.Mock, crawler: Crawler, spider: Spider
) -> None:
fixed_datetime = datetime(2019, 12, 1, 11, 38) fixed_datetime = datetime(2019, 12, 1, 11, 38)
mock_datetime.now = mock.Mock(return_value=fixed_datetime) mock_datetime.now = mock.Mock(return_value=fixed_datetime)
self.crawler.stats = StatsCollector(self.crawler) crawler.stats = StatsCollector(crawler)
ext = CoreStats.from_crawler(self.crawler) ext = CoreStats.from_crawler(crawler)
ext.spider_opened(self.spider) ext.spider_opened(spider)
ext.item_scraped({}, self.spider) ext.item_scraped({}, spider)
ext.response_received(self.spider) ext.response_received(spider)
ext.item_dropped({}, self.spider, ZeroDivisionError()) ext.item_dropped({}, spider, ZeroDivisionError())
ext.spider_closed(self.spider, "finished") ext.spider_closed(spider, "finished")
assert ext.stats._stats == { assert ext.stats._stats == {
"start_time": fixed_datetime, "start_time": fixed_datetime,
"finish_time": fixed_datetime, "finish_time": fixed_datetime,
@ -34,24 +50,22 @@ class TestCoreStatsExtension:
"elapsed_time_seconds": 0.0, "elapsed_time_seconds": 0.0,
} }
def test_core_stats_dummy_stats_collector(self): def test_core_stats_dummy_stats_collector(
self.crawler.stats = DummyStatsCollector(self.crawler) self, crawler: Crawler, spider: Spider
ext = CoreStats.from_crawler(self.crawler) ) -> None:
ext.spider_opened(self.spider) crawler.stats = DummyStatsCollector(crawler)
ext.item_scraped({}, self.spider) ext = CoreStats.from_crawler(crawler)
ext.response_received(self.spider) ext.spider_opened(spider)
ext.item_dropped({}, self.spider, ZeroDivisionError()) ext.item_scraped({}, spider)
ext.spider_closed(self.spider, "finished") ext.response_received(spider)
ext.item_dropped({}, spider, ZeroDivisionError())
ext.spider_closed(spider, "finished")
assert ext.stats._stats == {} assert ext.stats._stats == {}
class TestStatsCollector: class TestStatsCollector:
def setup_method(self): def test_collector(self, crawler: Crawler) -> None:
self.crawler = get_crawler(Spider) stats = StatsCollector(crawler)
self.spider = self.crawler._create_spider("foo")
def test_collector(self):
stats = StatsCollector(self.crawler)
assert stats.get_stats() == {} assert stats.get_stats() == {}
assert stats.get_value("anything") is None assert stats.get_value("anything") is None
assert stats.get_value("anything", "default") == "default" assert stats.get_value("anything", "default") == "default"
@ -77,8 +91,8 @@ class TestStatsCollector:
stats.min_value("test4", 7) stats.min_value("test4", 7)
assert stats.get_value("test4") == 7 assert stats.get_value("test4") == 7
def test_dummy_collector(self): def test_dummy_collector(self, crawler: Crawler, spider: Spider) -> None:
stats = DummyStatsCollector(self.crawler) stats = DummyStatsCollector(crawler)
assert stats.get_stats() == {} assert stats.get_stats() == {}
assert stats.get_value("anything") is None assert stats.get_value("anything") is None
assert stats.get_value("anything", "default") == "default" assert stats.get_value("anything", "default") == "default"
@ -86,7 +100,7 @@ class TestStatsCollector:
stats.inc_value("v1") stats.inc_value("v1")
stats.max_value("v2", 100) stats.max_value("v2", 100)
stats.min_value("v3", 100) stats.min_value("v3", 100)
stats.open_spider("a") stats.open_spider(spider)
stats.set_value("test", "value", spider=self.spider) stats.set_value("test", "value", spider=spider)
assert stats.get_stats() == {} assert stats.get_stats() == {}
assert stats.get_stats("a") == {} assert stats.get_stats(spider) == {}