Merge pull request #6130 from wRAR/typing-spider-mw

Full typing for scrapy/spidermiddlewares.
This commit is contained in:
Andrey Rakhmatullin 2023-11-03 12:31:19 +04:00 committed by GitHub
commit f3561807a6
No known key found for this signature in database
GPG Key ID: 4AEE18F83AFDEB23
5 changed files with 178 additions and 76 deletions

View File

@ -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

View File

@ -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

View File

@ -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):

View File

@ -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:

View File

@ -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
)