diff --git a/scrapy/spidermiddlewares/depth.py b/scrapy/spidermiddlewares/depth.py index eadc7c6ab..1e96654e2 100644 --- a/scrapy/spidermiddlewares/depth.py +++ b/scrapy/spidermiddlewares/depth.py @@ -4,46 +4,67 @@ Depth Spider Middleware See documentation in docs/topics/spider-middleware.rst """ -import logging +from __future__ import annotations -from scrapy.http import Request +import logging +from typing import TYPE_CHECKING, Any, AsyncIterable, Iterable + +from scrapy import Spider +from scrapy.crawler import Crawler +from scrapy.http import Request, Response +from scrapy.statscollectors import StatsCollector + +if TYPE_CHECKING: + # typing.Self requires Python 3.11 + from typing_extensions import Self logger = logging.getLogger(__name__) class DepthMiddleware: - def __init__(self, maxdepth, stats, verbose_stats=False, prio=1): + def __init__( + self, + maxdepth: int, + stats: StatsCollector, + verbose_stats: bool = False, + prio: int = 1, + ): self.maxdepth = maxdepth self.stats = stats self.verbose_stats = verbose_stats self.prio = prio @classmethod - def from_crawler(cls, crawler): + def from_crawler(cls, crawler: Crawler) -> Self: settings = crawler.settings maxdepth = settings.getint("DEPTH_LIMIT") verbose = settings.getbool("DEPTH_STATS_VERBOSE") prio = settings.getint("DEPTH_PRIORITY") + assert crawler.stats return cls(maxdepth, crawler.stats, verbose, prio) - def process_spider_output(self, response, result, spider): + def process_spider_output( + self, response: Response, result: Iterable[Any], spider: Spider + ) -> Iterable[Any]: self._init_depth(response, spider) - return (r for r in result or () if self._filter(r, response, spider)) + return (r for r in result if self._filter(r, response, spider)) - async def process_spider_output_async(self, response, result, spider): + async def process_spider_output_async( + self, response: Response, result: AsyncIterable[Any], spider: Spider + ) -> AsyncIterable[Any]: self._init_depth(response, spider) - async for r in result or (): + async for r in result: if self._filter(r, response, spider): yield r - def _init_depth(self, response, spider): + def _init_depth(self, response: Response, spider: Spider) -> None: # base case (depth=0) if "depth" not in response.meta: response.meta["depth"] = 0 if self.verbose_stats: self.stats.inc_value("request_depth_count/0", spider=spider) - def _filter(self, request, response, spider): + def _filter(self, request: Any, response: Response, spider: Spider) -> bool: if not isinstance(request, Request): return True depth = response.meta["depth"] + 1 diff --git a/scrapy/spidermiddlewares/httperror.py b/scrapy/spidermiddlewares/httperror.py index 0d3e5fe0b..94450b35b 100644 --- a/scrapy/spidermiddlewares/httperror.py +++ b/scrapy/spidermiddlewares/httperror.py @@ -3,9 +3,20 @@ HttpError Spider Middleware See documentation in docs/topics/spider-middleware.rst """ -import logging +from __future__ import annotations +import logging +from typing import TYPE_CHECKING, Any, Iterable, List, Optional + +from scrapy import Spider +from scrapy.crawler import Crawler from scrapy.exceptions import IgnoreRequest +from scrapy.http import Response +from scrapy.settings import BaseSettings + +if TYPE_CHECKING: + # typing.Self requires Python 3.11 + from typing_extensions import Self logger = logging.getLogger(__name__) @@ -13,21 +24,23 @@ logger = logging.getLogger(__name__) class HttpError(IgnoreRequest): """A non-200 response was filtered""" - def __init__(self, response, *args, **kwargs): + def __init__(self, response: Response, *args: Any, **kwargs: Any): self.response = response super().__init__(*args, **kwargs) class HttpErrorMiddleware: @classmethod - def from_crawler(cls, crawler): + def from_crawler(cls, crawler: Crawler) -> Self: return cls(crawler.settings) - def __init__(self, settings): - self.handle_httpstatus_all = settings.getbool("HTTPERROR_ALLOW_ALL") - self.handle_httpstatus_list = settings.getlist("HTTPERROR_ALLOWED_CODES") + def __init__(self, settings: BaseSettings): + self.handle_httpstatus_all: bool = settings.getbool("HTTPERROR_ALLOW_ALL") + self.handle_httpstatus_list: List[int] = settings.getlist( + "HTTPERROR_ALLOWED_CODES" + ) - def process_spider_input(self, response, spider): + def process_spider_input(self, response: Response, spider: Spider) -> None: if 200 <= response.status < 300: # common case return meta = response.meta @@ -45,8 +58,11 @@ class HttpErrorMiddleware: return raise HttpError(response, "Ignoring non-200 response") - def process_spider_exception(self, response, exception, spider): + def process_spider_exception( + self, response: Response, exception: Exception, spider: Spider + ) -> Optional[Iterable[Any]]: if isinstance(exception, HttpError): + assert spider.crawler.stats spider.crawler.stats.inc_value("httperror/response_ignored_count") spider.crawler.stats.inc_value( f"httperror/response_ignored_status_count/{response.status}" @@ -57,3 +73,4 @@ class HttpErrorMiddleware: extra={"spider": spider}, ) return [] + return None diff --git a/scrapy/spidermiddlewares/offsite.py b/scrapy/spidermiddlewares/offsite.py index 1a48926b3..a5214702d 100644 --- a/scrapy/spidermiddlewares/offsite.py +++ b/scrapy/spidermiddlewares/offsite.py @@ -3,36 +3,51 @@ Offsite Spider Middleware See documentation in docs/topics/spider-middleware.rst """ +from __future__ import annotations + import logging import re import warnings +from typing import TYPE_CHECKING, Any, AsyncIterable, Iterable, Set -from scrapy import signals -from scrapy.http import Request +from scrapy import Spider, signals +from scrapy.crawler import Crawler +from scrapy.http import Request, Response +from scrapy.statscollectors import StatsCollector from scrapy.utils.httpobj import urlparse_cached +if TYPE_CHECKING: + # typing.Self requires Python 3.11 + from typing_extensions import Self + + logger = logging.getLogger(__name__) class OffsiteMiddleware: - def __init__(self, stats): - self.stats = stats + def __init__(self, stats: StatsCollector): + self.stats: StatsCollector = stats @classmethod - def from_crawler(cls, crawler): + def from_crawler(cls, crawler: Crawler) -> Self: + assert crawler.stats o = cls(crawler.stats) crawler.signals.connect(o.spider_opened, signal=signals.spider_opened) return o - def process_spider_output(self, response, result, spider): - return (r for r in result or () if self._filter(r, spider)) + def process_spider_output( + self, response: Response, result: Iterable[Any], spider: Spider + ) -> Iterable[Any]: + return (r for r in result if self._filter(r, spider)) - async def process_spider_output_async(self, response, result, spider): - async for r in result or (): + 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, spider) -> bool: + def _filter(self, request: Any, spider: Spider) -> bool: if not isinstance(request, Request): return True if request.dont_filter or self.should_follow(request, spider): @@ -49,13 +64,13 @@ class OffsiteMiddleware: self.stats.inc_value("offsite/filtered", spider=spider) return False - def should_follow(self, request, spider): + def should_follow(self, request: Request, spider: Spider) -> bool: regex = self.host_regex # hostname can be None for wrong urls (like javascript links) host = urlparse_cached(request).hostname or "" return bool(regex.search(host)) - def get_host_regex(self, spider): + def get_host_regex(self, spider: Spider) -> re.Pattern[str]: """Override this method to implement a different offsite policy""" allowed_domains = getattr(spider, "allowed_domains", None) if not allowed_domains: @@ -83,9 +98,9 @@ class OffsiteMiddleware: regex = rf'^(.*\.)?({"|".join(domains)})$' return re.compile(regex) - def spider_opened(self, spider): - self.host_regex = self.get_host_regex(spider) - self.domains_seen = set() + def spider_opened(self, spider: Spider) -> None: + self.host_regex: re.Pattern[str] = self.get_host_regex(spider) + self.domains_seen: Set[str] = set() class URLWarning(Warning): diff --git a/scrapy/spidermiddlewares/referer.py b/scrapy/spidermiddlewares/referer.py index fd91e658b..a29e0ebb5 100644 --- a/scrapy/spidermiddlewares/referer.py +++ b/scrapy/spidermiddlewares/referer.py @@ -2,20 +2,39 @@ RefererMiddleware: populates Request referer field, based on the Response which originated it. """ +from __future__ import annotations + import warnings -from typing import Tuple +from typing import ( + TYPE_CHECKING, + Any, + AsyncIterable, + Dict, + Iterable, + Optional, + Tuple, + Type, + Union, + cast, +) from urllib.parse import urlparse from w3lib.url import safe_url_string -from scrapy import signals +from scrapy import Spider, signals +from scrapy.crawler import Crawler from scrapy.exceptions import NotConfigured from scrapy.http import Request, Response +from scrapy.settings import BaseSettings from scrapy.utils.misc import load_object from scrapy.utils.python import to_unicode from scrapy.utils.url import strip_url -LOCAL_SCHEMES = ( +if TYPE_CHECKING: + # typing.Self requires Python 3.11 + from typing_extensions import Self + +LOCAL_SCHEMES: Tuple[str, ...] = ( "about", "blob", "data", @@ -37,18 +56,20 @@ class ReferrerPolicy: NOREFERRER_SCHEMES: Tuple[str, ...] = LOCAL_SCHEMES name: str - def referrer(self, response_url, request_url): + def referrer(self, response_url: str, request_url: str) -> Optional[str]: raise NotImplementedError() - def stripped_referrer(self, url): + def stripped_referrer(self, url: str) -> Optional[str]: if urlparse(url).scheme not in self.NOREFERRER_SCHEMES: return self.strip_url(url) + return None - def origin_referrer(self, url): + def origin_referrer(self, url: str) -> Optional[str]: if urlparse(url).scheme not in self.NOREFERRER_SCHEMES: return self.origin(url) + return None - def strip_url(self, url, origin_only=False): + def strip_url(self, url: str, origin_only: bool = False) -> Optional[str]: """ https://www.w3.org/TR/referrer-policy/#strip-url @@ -72,18 +93,18 @@ class ReferrerPolicy: origin_only=origin_only, ) - def origin(self, url): + def origin(self, url: str) -> Optional[str]: """Return serialized origin (scheme, host, path) for a request or response URL.""" return self.strip_url(url, origin_only=True) - def potentially_trustworthy(self, url): + def potentially_trustworthy(self, url: str) -> bool: # Note: this does not follow https://w3c.github.io/webappsec-secure-contexts/#is-url-trustworthy parsed_url = urlparse(url) if parsed_url.scheme in ("data",): return False return self.tls_protected(url) - def tls_protected(self, url): + def tls_protected(self, url: str) -> bool: return urlparse(url).scheme in ("https", "ftps") @@ -98,7 +119,7 @@ class NoReferrerPolicy(ReferrerPolicy): name: str = POLICY_NO_REFERRER - def referrer(self, response_url, request_url): + def referrer(self, response_url: str, request_url: str) -> Optional[str]: return None @@ -119,9 +140,10 @@ class NoReferrerWhenDowngradePolicy(ReferrerPolicy): name: str = POLICY_NO_REFERRER_WHEN_DOWNGRADE - def referrer(self, response_url, request_url): + def referrer(self, response_url: str, request_url: str) -> Optional[str]: if not self.tls_protected(response_url) or self.tls_protected(request_url): return self.stripped_referrer(response_url) + return None class SameOriginPolicy(ReferrerPolicy): @@ -137,9 +159,10 @@ class SameOriginPolicy(ReferrerPolicy): name: str = POLICY_SAME_ORIGIN - def referrer(self, response_url, request_url): + def referrer(self, response_url: str, request_url: str) -> Optional[str]: if self.origin(response_url) == self.origin(request_url): return self.stripped_referrer(response_url) + return None class OriginPolicy(ReferrerPolicy): @@ -154,7 +177,7 @@ class OriginPolicy(ReferrerPolicy): name: str = POLICY_ORIGIN - def referrer(self, response_url, request_url): + def referrer(self, response_url: str, request_url: str) -> Optional[str]: return self.origin_referrer(response_url) @@ -174,13 +197,14 @@ class StrictOriginPolicy(ReferrerPolicy): name: str = POLICY_STRICT_ORIGIN - def referrer(self, response_url, request_url): + def referrer(self, response_url: str, request_url: str) -> Optional[str]: if ( self.tls_protected(response_url) and self.potentially_trustworthy(request_url) or not self.tls_protected(response_url) ): return self.origin_referrer(response_url) + return None class OriginWhenCrossOriginPolicy(ReferrerPolicy): @@ -197,7 +221,7 @@ class OriginWhenCrossOriginPolicy(ReferrerPolicy): name: str = POLICY_ORIGIN_WHEN_CROSS_ORIGIN - def referrer(self, response_url, request_url): + def referrer(self, response_url: str, request_url: str) -> Optional[str]: origin = self.origin(response_url) if origin == self.origin(request_url): return self.stripped_referrer(response_url) @@ -224,7 +248,7 @@ class StrictOriginWhenCrossOriginPolicy(ReferrerPolicy): name: str = POLICY_STRICT_ORIGIN_WHEN_CROSS_ORIGIN - def referrer(self, response_url, request_url): + def referrer(self, response_url: str, request_url: str) -> Optional[str]: origin = self.origin(response_url) if origin == self.origin(request_url): return self.stripped_referrer(response_url) @@ -234,6 +258,7 @@ class StrictOriginWhenCrossOriginPolicy(ReferrerPolicy): or not self.tls_protected(response_url) ): return self.origin_referrer(response_url) + return None class UnsafeUrlPolicy(ReferrerPolicy): @@ -252,7 +277,7 @@ class UnsafeUrlPolicy(ReferrerPolicy): name: str = POLICY_UNSAFE_URL - def referrer(self, response_url, request_url): + def referrer(self, response_url: str, request_url: str) -> Optional[str]: return self.stripped_referrer(response_url) @@ -267,7 +292,7 @@ class DefaultReferrerPolicy(NoReferrerWhenDowngradePolicy): name: str = POLICY_SCRAPY_DEFAULT -_policy_classes = { +_policy_classes: Dict[str, Type[ReferrerPolicy]] = { p.name: p for p in ( NoReferrerPolicy, @@ -286,14 +311,16 @@ _policy_classes = { _policy_classes[""] = NoReferrerWhenDowngradePolicy -def _load_policy_class(policy, warning_only=False): +def _load_policy_class( + policy: str, warning_only: bool = False +) -> Optional[Type[ReferrerPolicy]]: """ Expect a string for the path to the policy class, otherwise try to interpret the string as a standard value from https://www.w3.org/TR/referrer-policy/#referrer-policies """ try: - return load_object(policy) + return cast(Type[ReferrerPolicy], load_object(policy)) except ValueError: try: return _policy_classes[policy.lower()] @@ -307,13 +334,15 @@ def _load_policy_class(policy, warning_only=False): class RefererMiddleware: - def __init__(self, settings=None): - self.default_policy = DefaultReferrerPolicy + def __init__(self, settings: Optional[BaseSettings] = None): + self.default_policy: Type[ReferrerPolicy] = DefaultReferrerPolicy if settings is not None: - self.default_policy = _load_policy_class(settings.get("REFERRER_POLICY")) + settings_policy = _load_policy_class(settings.get("REFERRER_POLICY")) + assert settings_policy + self.default_policy = settings_policy @classmethod - def from_crawler(cls, crawler): + def from_crawler(cls, crawler: Crawler) -> Self: if not crawler.settings.getbool("REFERER_ENABLED"): raise NotConfigured mw = cls(crawler.settings) @@ -323,7 +352,9 @@ class RefererMiddleware: return mw - def policy(self, resp_or_url, request): + def policy( + self, resp_or_url: Union[Response, str], request: Request + ) -> ReferrerPolicy: """ Determine Referrer-Policy to use from a parent Response (or URL), and a Request to be sent. @@ -348,21 +379,25 @@ class RefererMiddleware: cls = _load_policy_class(policy_name, warning_only=True) return cls() if cls else self.default_policy() - def process_spider_output(self, response, result, spider): - return (self._set_referer(r, response) for r in result or ()) + def process_spider_output( + self, response: Response, result: Iterable[Any], spider: Spider + ) -> Iterable[Any]: + return (self._set_referer(r, response) for r in result) - async def process_spider_output_async(self, response, result, spider): - async for r in result or (): + async def process_spider_output_async( + self, response: Response, result: AsyncIterable[Any], spider: Spider + ) -> AsyncIterable[Any]: + async for r in result: yield self._set_referer(r, response) - def _set_referer(self, 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, spider): + def request_scheduled(self, request: Request, spider: Spider) -> None: # check redirected request to patch "Referer" header if necessary redirected_urls = request.meta.get("redirect_urls", []) if redirected_urls: @@ -378,7 +413,7 @@ class RefererMiddleware: policy_referrer = self.policy(parent_url, request).referrer( parent_url, request.url ) - if policy_referrer != request_referrer: + if policy_referrer != request_referrer.decode("latin1"): if policy_referrer is None: request.headers.pop("Referer") else: diff --git a/scrapy/spidermiddlewares/urllength.py b/scrapy/spidermiddlewares/urllength.py index f6d92e53a..e2aa554a7 100644 --- a/scrapy/spidermiddlewares/urllength.py +++ b/scrapy/spidermiddlewares/urllength.py @@ -4,40 +4,54 @@ Url Length Spider Middleware See documentation in docs/topics/spider-middleware.rst """ -import logging +from __future__ import annotations +import logging +from typing import TYPE_CHECKING, Any, AsyncIterable, Iterable + +from scrapy import Spider from scrapy.exceptions import NotConfigured -from scrapy.http import Request +from scrapy.http import Request, Response +from scrapy.settings import BaseSettings + +if TYPE_CHECKING: + # typing.Self requires Python 3.11 + from typing_extensions import Self logger = logging.getLogger(__name__) class UrlLengthMiddleware: - def __init__(self, maxlength): - self.maxlength = maxlength + def __init__(self, maxlength: int): + self.maxlength: int = maxlength @classmethod - def from_settings(cls, settings): + def from_settings(cls, settings: BaseSettings) -> Self: maxlength = settings.getint("URLLENGTH_LIMIT") if not maxlength: raise NotConfigured return cls(maxlength) - def process_spider_output(self, response, result, spider): - return (r for r in result or () if self._filter(r, spider)) + def process_spider_output( + self, response: Response, result: Iterable[Any], spider: Spider + ) -> Iterable[Any]: + return (r for r in result if self._filter(r, spider)) - async def process_spider_output_async(self, response, result, spider): - async for r in result or (): + 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, spider): + def _filter(self, request: Any, spider: Spider) -> bool: if isinstance(request, Request) and len(request.url) > self.maxlength: logger.info( "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 )