Base class for universal spider middlewares (#6693)

* Initial BaseSpiderMiddleware.

* Rename the new methods.

* Remove the spider argument from new BaseSpiderMiddleware methods.

* Add docs for BaseSpiderMiddleware.

* Silence pylint.

* Add BaseSpiderMiddleware tests.

* Add a release note.
This commit is contained in:
Andrey Rakhmatullin 2025-04-23 18:29:04 +04:00 committed by GitHub
parent 9f99da8f86
commit daf9db72b2
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
10 changed files with 329 additions and 120 deletions

View File

@ -3,6 +3,23 @@
Release notes Release notes
============= =============
.. _release-VERSION:
Scrapy VERSION (unreleased)
---------------------------
Backward-incompatible changes
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
- The ``from_settings()`` method of
:class:`~scrapy.spidermiddlewares.urllength.UrlLengthMiddleware` is removed
without a deprecation period (this was needed because after the
introduction of the
:class:`~scrapy.spidermiddlewares.base.BaseSpiderMiddleware` base class and
switching built-in spider middlewares to it those middlewares need the
:class:`~scrapy.crawler.Crawler` instance at run time). Please use
``from_crawler()`` instead.
.. _release-2.12.0: .. _release-2.12.0:
Scrapy 2.12.0 (2024-11-18) Scrapy 2.12.0 (2024-11-18)

View File

@ -189,6 +189,19 @@ one or more of these methods:
:param spider: the spider which raised the exception :param spider: the spider which raised the exception
:type spider: :class:`~scrapy.Spider` object :type spider: :class:`~scrapy.Spider` object
Base class for custom spider middlewares
----------------------------------------
Scrapy provides a base class for custom spider middlewares. It's not required
to use it but it can help with simplifying middleware implementations and
reducing the amount of boilerplate code in :ref:`universal middlewares
<universal-spider-middleware>`.
.. module:: scrapy.spidermiddlewares.base
.. autoclass:: BaseSpiderMiddleware
:members:
.. _topics-spider-middleware-ref: .. _topics-spider-middleware-ref:
Built-in spider middleware reference Built-in spider middleware reference

View File

@ -0,0 +1,97 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any
from scrapy import Request, Spider
if TYPE_CHECKING:
from collections.abc import AsyncIterable, Iterable
# typing.Self requires Python 3.11
from typing_extensions import Self
from scrapy.crawler import Crawler
from scrapy.http import Response
class BaseSpiderMiddleware:
"""Optional base class for spider middlewares.
This class provides helper methods for asynchronous ``process_spider_output``
methods. Middlewares that don't have a ``process_spider_output`` method don't need
to use it.
You can override the
:meth:`~scrapy.spidermiddlewares.base.BaseSpiderMiddleware.get_processed_request`
method to add processing code for requests and the
:meth:`~scrapy.spidermiddlewares.base.BaseSpiderMiddleware.get_processed_item`
method to add processing code for items. These methods take a single
request or item from the spider output iterable and return a request or
item (the same or a new one), or ``None`` to remove this request or item
from the processing.
"""
def __init__(self, crawler: Crawler):
self.crawler: Crawler = crawler
@classmethod
def from_crawler(cls, crawler: Crawler) -> Self:
return cls(crawler)
def process_spider_output(
self, response: Response, result: Iterable[Any], spider: Spider
) -> Iterable[Any]:
for o in result:
if isinstance(o, Request):
o = self.get_processed_request(o, response)
else:
o = self.get_processed_item(o, response)
if o is not None:
yield o
async def process_spider_output_async(
self, response: Response, result: AsyncIterable[Any], spider: Spider
) -> AsyncIterable[Any]:
async for o in result:
if isinstance(o, Request):
o = self.get_processed_request(o, response)
else:
o = self.get_processed_item(o, response)
if o is not None:
yield o
def get_processed_request(
self, request: Request, response: Response
) -> Request | None:
"""Return a processed request from the spider output.
This method is called with a single request from the spider output.
It should return the same or a different request, or ``None`` to
ignore it.
:param request: the input request
:type request: :class:`~scrapy.Request` object
:param response: the response being processed
:type response: :class:`~scrapy.http.Response` object
:return: the processed request or ``None``
"""
return request
def get_processed_item(self, item: Any, response: Response) -> Any:
"""Return a processed item from the spider output.
This method is called with a single item from the spider output.
It should return the same or a different item, or ``None`` to
ignore it.
:param item: the input item
:type item: item object
:param response: the response being processed
:type response: :class:`~scrapy.http.Response` object
:return: the processed item or ``None``
"""
return item

View File

@ -9,7 +9,7 @@ from __future__ import annotations
import logging import logging
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from scrapy.http import Request, Response from scrapy.spidermiddlewares.base import BaseSpiderMiddleware
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncIterable, Iterable from collections.abc import AsyncIterable, Iterable
@ -19,14 +19,17 @@ if TYPE_CHECKING:
from scrapy import Spider from scrapy import Spider
from scrapy.crawler import Crawler from scrapy.crawler import Crawler
from scrapy.http import Request, Response
from scrapy.statscollectors import StatsCollector from scrapy.statscollectors import StatsCollector
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class DepthMiddleware: class DepthMiddleware(BaseSpiderMiddleware):
def __init__( crawler: Crawler
def __init__( # pylint: disable=super-init-not-called
self, self,
maxdepth: int, maxdepth: int,
stats: StatsCollector, stats: StatsCollector,
@ -45,21 +48,22 @@ class DepthMiddleware:
verbose = settings.getbool("DEPTH_STATS_VERBOSE") verbose = settings.getbool("DEPTH_STATS_VERBOSE")
prio = settings.getint("DEPTH_PRIORITY") prio = settings.getint("DEPTH_PRIORITY")
assert crawler.stats assert crawler.stats
return cls(maxdepth, crawler.stats, verbose, prio) o = cls(maxdepth, crawler.stats, verbose, prio)
o.crawler = crawler
return o
def process_spider_output( def process_spider_output(
self, response: Response, result: Iterable[Any], spider: Spider self, response: Response, result: Iterable[Any], spider: Spider
) -> Iterable[Any]: ) -> Iterable[Any]:
self._init_depth(response, spider) self._init_depth(response, spider)
return (r for r in result if self._filter(r, response, spider)) yield from super().process_spider_output(response, result, spider)
async def process_spider_output_async( async def process_spider_output_async(
self, response: Response, result: AsyncIterable[Any], spider: Spider self, response: Response, result: AsyncIterable[Any], spider: Spider
) -> AsyncIterable[Any]: ) -> AsyncIterable[Any]:
self._init_depth(response, spider) self._init_depth(response, spider)
async for r in result: async for o in super().process_spider_output_async(response, result, spider):
if self._filter(r, response, spider): yield o
yield r
def _init_depth(self, response: Response, spider: Spider) -> None: def _init_depth(self, response: Response, spider: Spider) -> None:
# base case (depth=0) # base case (depth=0)
@ -68,9 +72,9 @@ class DepthMiddleware:
if self.verbose_stats: if self.verbose_stats:
self.stats.inc_value("request_depth_count/0", spider=spider) self.stats.inc_value("request_depth_count/0", spider=spider)
def _filter(self, request: Any, response: Response, spider: Spider) -> bool: def get_processed_request(
if not isinstance(request, Request): self, request: Request, response: Response
return True ) -> Request | None:
depth = response.meta["depth"] + 1 depth = response.meta["depth"] + 1
request.meta["depth"] = depth request.meta["depth"] = depth
if self.prio: if self.prio:
@ -79,10 +83,12 @@ class DepthMiddleware:
logger.debug( logger.debug(
"Ignoring link (depth > %(maxdepth)d): %(requrl)s ", "Ignoring link (depth > %(maxdepth)d): %(requrl)s ",
{"maxdepth": self.maxdepth, "requrl": request.url}, {"maxdepth": self.maxdepth, "requrl": request.url},
extra={"spider": spider}, extra={"spider": self.crawler.spider},
) )
return False return None
if self.verbose_stats: if self.verbose_stats:
self.stats.inc_value(f"request_depth_count/{depth}", spider=spider) self.stats.inc_value(
self.stats.max_value("request_depth_max", depth, spider=spider) f"request_depth_count/{depth}", spider=self.crawler.spider
return True )
self.stats.max_value("request_depth_max", depth, spider=self.crawler.spider)
return request

View File

@ -9,11 +9,11 @@ from __future__ import annotations
import logging import logging
import re import re
import warnings import warnings
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING
from scrapy import Spider, signals from scrapy import Spider, signals
from scrapy.exceptions import ScrapyDeprecationWarning from scrapy.exceptions import ScrapyDeprecationWarning
from scrapy.http import Request, Response from scrapy.spidermiddlewares.base import BaseSpiderMiddleware
from scrapy.utils.httpobj import urlparse_cached from scrapy.utils.httpobj import urlparse_cached
warnings.warn( warnings.warn(
@ -23,61 +23,52 @@ warnings.warn(
) )
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncIterable, Iterable
# typing.Self requires Python 3.11 # typing.Self requires Python 3.11
from typing_extensions import Self from typing_extensions import Self
from scrapy.crawler import Crawler from scrapy.crawler import Crawler
from scrapy.http import Request, Response
from scrapy.statscollectors import StatsCollector from scrapy.statscollectors import StatsCollector
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class OffsiteMiddleware: class OffsiteMiddleware(BaseSpiderMiddleware):
def __init__(self, stats: StatsCollector): crawler: Crawler
def __init__(self, stats: StatsCollector): # pylint: disable=super-init-not-called
self.stats: StatsCollector = stats self.stats: StatsCollector = stats
@classmethod @classmethod
def from_crawler(cls, crawler: Crawler) -> Self: def from_crawler(cls, crawler: Crawler) -> Self:
assert crawler.stats assert crawler.stats
o = cls(crawler.stats) o = cls(crawler.stats)
o.crawler = crawler
crawler.signals.connect(o.spider_opened, signal=signals.spider_opened) crawler.signals.connect(o.spider_opened, signal=signals.spider_opened)
return o return o
def process_spider_output( def get_processed_request(
self, response: Response, result: Iterable[Any], spider: Spider self, request: Request, response: Response
) -> Iterable[Any]: ) -> Request | None:
return (r for r in result if self._filter(r, spider)) assert self.crawler.spider
async def process_spider_output_async(
self, response: Response, result: AsyncIterable[Any], spider: Spider
) -> AsyncIterable[Any]:
async for r in result:
if self._filter(r, spider):
yield r
def _filter(self, request: Any, spider: Spider) -> bool:
if not isinstance(request, Request):
return True
if ( if (
request.dont_filter request.dont_filter
or request.meta.get("allow_offsite") or request.meta.get("allow_offsite")
or self.should_follow(request, spider) or self.should_follow(request, self.crawler.spider)
): ):
return True return request
domain = urlparse_cached(request).hostname domain = urlparse_cached(request).hostname
if domain and domain not in self.domains_seen: if domain and domain not in self.domains_seen:
self.domains_seen.add(domain) self.domains_seen.add(domain)
logger.debug( logger.debug(
"Filtered offsite request to %(domain)r: %(request)s", "Filtered offsite request to %(domain)r: %(request)s",
{"domain": domain, "request": request}, {"domain": domain, "request": request},
extra={"spider": spider}, extra={"spider": self.crawler.spider},
) )
self.stats.inc_value("offsite/domains", spider=spider) self.stats.inc_value("offsite/domains", spider=self.crawler.spider)
self.stats.inc_value("offsite/filtered", spider=spider) self.stats.inc_value("offsite/filtered", spider=self.crawler.spider)
return False return None
def should_follow(self, request: Request, spider: Spider) -> bool: def should_follow(self, request: Request, spider: Spider) -> bool:
regex = self.host_regex regex = self.host_regex

View File

@ -6,7 +6,7 @@ originated it.
from __future__ import annotations from __future__ import annotations
import warnings import warnings
from typing import TYPE_CHECKING, Any, cast from typing import TYPE_CHECKING, cast
from urllib.parse import urlparse from urllib.parse import urlparse
from w3lib.url import safe_url_string from w3lib.url import safe_url_string
@ -14,13 +14,12 @@ from w3lib.url import safe_url_string
from scrapy import Spider, signals from scrapy import Spider, signals
from scrapy.exceptions import NotConfigured from scrapy.exceptions import NotConfigured
from scrapy.http import Request, Response from scrapy.http import Request, Response
from scrapy.spidermiddlewares.base import BaseSpiderMiddleware
from scrapy.utils.misc import load_object from scrapy.utils.misc import load_object
from scrapy.utils.python import to_unicode from scrapy.utils.python import to_unicode
from scrapy.utils.url import strip_url from scrapy.utils.url import strip_url
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncIterable, Iterable
# typing.Self requires Python 3.11 # typing.Self requires Python 3.11
from typing_extensions import Self from typing_extensions import Self
@ -327,8 +326,8 @@ def _load_policy_class(
return None return None
class RefererMiddleware: class RefererMiddleware(BaseSpiderMiddleware):
def __init__(self, settings: BaseSettings | None = None): def __init__(self, settings: BaseSettings | None = None): # pylint: disable=super-init-not-called
self.default_policy: type[ReferrerPolicy] = DefaultReferrerPolicy self.default_policy: type[ReferrerPolicy] = DefaultReferrerPolicy
if settings is not None: if settings is not None:
settings_policy = _load_policy_class(settings.get("REFERRER_POLICY")) settings_policy = _load_policy_class(settings.get("REFERRER_POLICY"))
@ -370,23 +369,13 @@ class RefererMiddleware:
cls = _load_policy_class(policy_name, warning_only=True) cls = _load_policy_class(policy_name, warning_only=True)
return cls() if cls else self.default_policy() return cls() if cls else self.default_policy()
def process_spider_output( def get_processed_request(
self, response: Response, result: Iterable[Any], spider: Spider self, request: Request, response: Response
) -> Iterable[Any]: ) -> Request | None:
return (self._set_referer(r, response) for r in result) referrer = self.policy(response, request).referrer(response.url, request.url)
if referrer is not None:
async def process_spider_output_async( request.headers.setdefault("Referer", referrer)
self, response: Response, result: AsyncIterable[Any], spider: Spider return request
) -> AsyncIterable[Any]:
async for r in result:
yield self._set_referer(r, response)
def _set_referer(self, r: Any, response: Response) -> Any:
if isinstance(r, Request):
referrer = self.policy(response, r).referrer(response.url, r.url)
if referrer is not None:
r.headers.setdefault("Referer", referrer)
return r
def request_scheduled(self, request: Request, spider: Spider) -> None: def request_scheduled(self, request: Request, spider: Spider) -> None:
# check redirected request to patch "Referer" header if necessary # check redirected request to patch "Referer" header if necessary

View File

@ -7,72 +7,49 @@ See documentation in docs/topics/spider-middleware.rst
from __future__ import annotations from __future__ import annotations
import logging import logging
import warnings from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any
from scrapy.exceptions import NotConfigured, ScrapyDeprecationWarning from scrapy.exceptions import NotConfigured
from scrapy.http import Request, Response from scrapy.spidermiddlewares.base import BaseSpiderMiddleware
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncIterable, Iterable
# typing.Self requires Python 3.11 # typing.Self requires Python 3.11
from typing_extensions import Self from typing_extensions import Self
from scrapy import Spider
from scrapy.crawler import Crawler from scrapy.crawler import Crawler
from scrapy.settings import BaseSettings from scrapy.http import Request, Response
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class UrlLengthMiddleware: class UrlLengthMiddleware(BaseSpiderMiddleware):
def __init__(self, maxlength: int): crawler: Crawler
def __init__(self, maxlength: int): # pylint: disable=super-init-not-called
self.maxlength: int = maxlength self.maxlength: int = maxlength
@classmethod
def from_settings(cls, settings: BaseSettings) -> Self:
warnings.warn(
f"{cls.__name__}.from_settings() is deprecated, use from_crawler() instead.",
category=ScrapyDeprecationWarning,
stacklevel=2,
)
return cls._from_settings(settings)
@classmethod @classmethod
def from_crawler(cls, crawler: Crawler) -> Self: def from_crawler(cls, crawler: Crawler) -> Self:
return cls._from_settings(crawler.settings) maxlength = crawler.settings.getint("URLLENGTH_LIMIT")
@classmethod
def _from_settings(cls, settings: BaseSettings) -> Self:
maxlength = settings.getint("URLLENGTH_LIMIT")
if not maxlength: if not maxlength:
raise NotConfigured raise NotConfigured
return cls(maxlength) o = cls(maxlength)
o.crawler = crawler
return o
def process_spider_output( def get_processed_request(
self, response: Response, result: Iterable[Any], spider: Spider self, request: Request, response: Response
) -> Iterable[Any]: ) -> Request | None:
return (r for r in result if self._filter(r, spider)) if len(request.url) <= self.maxlength:
return request
async def process_spider_output_async( logger.info(
self, response: Response, result: AsyncIterable[Any], spider: Spider "Ignoring link (url length > %(maxlength)d): %(url)s ",
) -> AsyncIterable[Any]: {"maxlength": self.maxlength, "url": request.url},
async for r in result: extra={"spider": self.crawler.spider},
if self._filter(r, spider): )
yield r assert self.crawler.stats
self.crawler.stats.inc_value(
def _filter(self, request: Any, spider: Spider) -> bool: "urllength/request_ignored_count", spider=self.crawler.spider
if isinstance(request, Request) and len(request.url) > self.maxlength: )
logger.info( return None
"Ignoring link (url length > %(maxlength)d): %(url)s ",
{"maxlength": self.maxlength, "url": request.url},
extra={"spider": spider},
)
assert spider.crawler.stats
spider.crawler.stats.inc_value(
"urllength/request_ignored_count", spider=spider
)
return False
return True

View File

@ -0,0 +1,120 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any
import pytest
from scrapy import Request, Spider
from scrapy.http import Response
from scrapy.spidermiddlewares.base import BaseSpiderMiddleware
from scrapy.utils.test import get_crawler
if TYPE_CHECKING:
from scrapy.crawler import Crawler
@pytest.fixture
def crawler() -> Crawler:
return get_crawler(Spider)
def test_trivial(crawler):
class TrivialSpiderMiddleware(BaseSpiderMiddleware):
pass
mw = TrivialSpiderMiddleware.from_crawler(crawler)
assert hasattr(mw, "crawler")
assert mw.crawler is crawler
test_req = Request("data:,")
spider_output = [test_req, {"foo": "bar"}]
processed = list(
mw.process_spider_output(Response("data:,"), spider_output, crawler.spider)
)
assert processed == [test_req, {"foo": "bar"}]
def test_processed_request(crawler):
class ProcessReqSpiderMiddleware(BaseSpiderMiddleware):
def get_processed_request(
self, request: Request, response: Response
) -> Request | None:
if request.url == "data:2,":
return None
if request.url == "data:3,":
return Request("data:30,")
return request
mw = ProcessReqSpiderMiddleware.from_crawler(crawler)
test_req1 = Request("data:1,")
test_req2 = Request("data:2,")
test_req3 = Request("data:3,")
spider_output = [test_req1, {"foo": "bar"}, test_req2, test_req3]
processed = list(
mw.process_spider_output(Response("data:,"), spider_output, crawler.spider)
)
assert len(processed) == 3
assert isinstance(processed[0], Request)
assert processed[0].url == "data:1,"
assert processed[1] == {"foo": "bar"}
assert isinstance(processed[2], Request)
assert processed[2].url == "data:30,"
def test_processed_item(crawler):
class ProcessItemSpiderMiddleware(BaseSpiderMiddleware):
def get_processed_item(self, item: Any, response: Response) -> Any:
if item["foo"] == 2:
return None
if item["foo"] == 3:
item["foo"] = 30
return item
mw = ProcessItemSpiderMiddleware.from_crawler(crawler)
test_req = Request("data:,")
spider_output = [{"foo": 1}, {"foo": 2}, test_req, {"foo": 3}]
processed = list(
mw.process_spider_output(Response("data:,"), spider_output, crawler.spider)
)
assert processed == [{"foo": 1}, test_req, {"foo": 30}]
def test_processed_both(crawler):
class ProcessBothSpiderMiddleware(BaseSpiderMiddleware):
def get_processed_request(
self, request: Request, response: Response
) -> Request | None:
if request.url == "data:2,":
return None
if request.url == "data:3,":
return Request("data:30,")
return request
def get_processed_item(self, item: Any, response: Response) -> Any:
if item["foo"] == 2:
return None
if item["foo"] == 3:
item["foo"] = 30
return item
mw = ProcessBothSpiderMiddleware.from_crawler(crawler)
test_req1 = Request("data:1,")
test_req2 = Request("data:2,")
test_req3 = Request("data:3,")
spider_output = [
test_req1,
{"foo": 1},
{"foo": 2},
test_req2,
{"foo": 3},
test_req3,
]
processed = list(
mw.process_spider_output(Response("data:,"), spider_output, crawler.spider)
)
assert len(processed) == 4
assert isinstance(processed[0], Request)
assert processed[0].url == "data:1,"
assert processed[1] == {"foo": 1}
assert processed[2] == {"foo": 30}
assert isinstance(processed[3], Request)
assert processed[3].url == "data:30,"

View File

@ -1,19 +1,18 @@
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.statscollectors import StatsCollector
from scrapy.utils.test import get_crawler from scrapy.utils.test import get_crawler
class TestDepthMiddleware: class TestDepthMiddleware:
def setup_method(self): def setup_method(self):
crawler = get_crawler(Spider) crawler = get_crawler(Spider, {"DEPTH_LIMIT": 1, "DEPTH_STATS_VERBOSE": True})
self.spider = crawler._create_spider("scrapytest.org") self.spider = crawler._create_spider("scrapytest.org")
self.stats = StatsCollector(crawler) self.stats = crawler.stats
self.stats.open_spider(self.spider) self.stats.open_spider(self.spider)
self.mw = DepthMiddleware(1, self.stats, True) self.mw = DepthMiddleware.from_crawler(crawler)
def test_process_spider_output(self): def test_process_spider_output(self):
req = Request("http://scrapytest.org") req = Request("http://scrapytest.org")

View File

@ -10,7 +10,7 @@ from scrapy.utils.test import get_crawler
class TestOffsiteMiddleware: class TestOffsiteMiddleware:
def setup_method(self): def setup_method(self):
crawler = get_crawler(Spider) crawler = get_crawler(Spider)
self.spider = crawler._create_spider(**self._get_spiderargs()) self.spider = crawler.spider = crawler._create_spider(**self._get_spiderargs())
self.mw = OffsiteMiddleware.from_crawler(crawler) self.mw = OffsiteMiddleware.from_crawler(crawler)
self.mw.spider_opened(self.spider) self.mw.spider_opened(self.spider)