From 8a08283580176049cb539423795830b9faea91a9 Mon Sep 17 00:00:00 2001 From: Andrey Rakhmatullin Date: Sun, 5 May 2024 22:32:46 +0500 Subject: [PATCH] Full typing for scrapy/http/cookies.py. --- scrapy/http/cookies.py | 123 +++++++++++++++++++++-------------- scrapy/http/response/text.py | 4 +- 2 files changed, 76 insertions(+), 51 deletions(-) diff --git a/scrapy/http/cookies.py b/scrapy/http/cookies.py index 72855bad5..8af89c74f 100644 --- a/scrapy/http/cookies.py +++ b/scrapy/http/cookies.py @@ -1,36 +1,56 @@ +from __future__ import annotations + import re import time from http.cookiejar import Cookie from http.cookiejar import CookieJar as _CookieJar -from http.cookiejar import DefaultCookiePolicy -from typing import Sequence +from http.cookiejar import CookiePolicy, DefaultCookiePolicy +from typing import ( + TYPE_CHECKING, + Any, + Dict, + Iterator, + List, + Optional, + Sequence, + Tuple, + cast, +) from scrapy import Request from scrapy.http import Response from scrapy.utils.httpobj import urlparse_cached from scrapy.utils.python import to_unicode +if TYPE_CHECKING: + # typing.Self requires Python 3.11 + from typing_extensions import Self + # Defined in the http.cookiejar module, but undocumented: # https://github.com/python/cpython/blob/v3.9.0/Lib/http/cookiejar.py#L527 IPV4_RE = re.compile(r"\.\d+$", re.ASCII) class CookieJar: - def __init__(self, policy=None, check_expired_frequency=10000): - self.policy = policy or DefaultCookiePolicy() - self.jar = _CookieJar(self.policy) - self.jar._cookies_lock = _DummyLock() - self.check_expired_frequency = check_expired_frequency - self.processed = 0 + def __init__( + self, + policy: Optional[CookiePolicy] = None, + check_expired_frequency: int = 10000, + ): + self.policy: CookiePolicy = policy or DefaultCookiePolicy() + self.jar: _CookieJar = _CookieJar(self.policy) + self.jar._cookies_lock = _DummyLock() # type: ignore[attr-defined] + self.check_expired_frequency: int = check_expired_frequency + self.processed: int = 0 - def extract_cookies(self, response, request): + def extract_cookies(self, response: Response, request: Request) -> None: wreq = WrappedRequest(request) wrsp = WrappedResponse(response) - return self.jar.extract_cookies(wrsp, wreq) + self.jar.extract_cookies(wrsp, wreq) # type: ignore[arg-type] def add_cookie_header(self, request: Request) -> None: wreq = WrappedRequest(request) - self.policy._now = self.jar._now = int(time.time()) + self.policy._now = self.jar._now = int(time.time()) # type: ignore[attr-defined] # the cookiejar implementation iterates through all domains # instead we restrict to potential matches on the domain @@ -47,10 +67,10 @@ class CookieJar: cookies = [] for host in hosts: - if host in self.jar._cookies: - cookies += self.jar._cookies_for_domain(host, wreq) + if host in self.jar._cookies: # type: ignore[attr-defined] + cookies += self.jar._cookies_for_domain(host, wreq) # type: ignore[attr-defined] - attrs = self.jar._cookie_attrs(cookies) + attrs = self.jar._cookie_attrs(cookies) # type: ignore[attr-defined] if attrs: if not wreq.has_header("Cookie"): wreq.add_unredirected_header("Cookie", "; ".join(attrs)) @@ -61,37 +81,42 @@ class CookieJar: self.jar.clear_expired_cookies() @property - def _cookies(self): - return self.jar._cookies + def _cookies(self) -> Dict[str, Dict[str, Dict[str, Cookie]]]: + return self.jar._cookies # type: ignore[attr-defined,no-any-return] - def clear_session_cookies(self, *args, **kwargs): - return self.jar.clear_session_cookies(*args, **kwargs) + def clear_session_cookies(self) -> None: + return self.jar.clear_session_cookies() - def clear(self, domain=None, path=None, name=None): - return self.jar.clear(domain, path, name) + def clear( + self, + domain: Optional[str] = None, + path: Optional[str] = None, + name: Optional[str] = None, + ) -> None: + self.jar.clear(domain, path, name) - def __iter__(self): + def __iter__(self) -> Iterator[Cookie]: return iter(self.jar) - def __len__(self): + def __len__(self) -> int: return len(self.jar) - def set_policy(self, pol): - return self.jar.set_policy(pol) + def set_policy(self, pol: CookiePolicy) -> None: + self.jar.set_policy(pol) def make_cookies(self, response: Response, request: Request) -> Sequence[Cookie]: wreq = WrappedRequest(request) wrsp = WrappedResponse(response) - return self.jar.make_cookies(wrsp, wreq) + return self.jar.make_cookies(wrsp, wreq) # type: ignore[arg-type] - def set_cookie(self, cookie): + def set_cookie(self, cookie: Cookie) -> None: self.jar.set_cookie(cookie) def set_cookie_if_ok(self, cookie: Cookie, request: Request) -> None: - self.jar.set_cookie_if_ok(cookie, WrappedRequest(request)) + self.jar.set_cookie_if_ok(cookie, WrappedRequest(request)) # type: ignore[arg-type] -def potential_domain_matches(domain): +def potential_domain_matches(domain: str) -> List[str]: """Potential domain matches for a cookie >>> potential_domain_matches('www.example.com') @@ -111,10 +136,10 @@ def potential_domain_matches(domain): class _DummyLock: - def acquire(self): + def acquire(self) -> None: pass - def release(self): + def release(self) -> None: pass @@ -124,19 +149,19 @@ class WrappedRequest: see http://docs.python.org/library/urllib2.html#urllib2.Request """ - def __init__(self, request): + def __init__(self, request: Request): self.request = request - def get_full_url(self): + def get_full_url(self) -> str: return self.request.url - def get_host(self): + def get_host(self) -> str: return urlparse_cached(self.request).netloc - def get_type(self): + def get_type(self) -> str: return urlparse_cached(self.request).scheme - def is_unverifiable(self): + def is_unverifiable(self) -> bool: """Unverifiable should indicate whether the request is unverifiable, as defined by RFC 2965. It defaults to False. An unverifiable request is one whose URL the user did not have the @@ -144,36 +169,36 @@ class WrappedRequest: HTML document, and the user had no option to approve the automatic fetching of the image, this should be true. """ - return self.request.meta.get("is_unverifiable", False) + return cast(bool, self.request.meta.get("is_unverifiable", False)) @property - def full_url(self): + def full_url(self) -> str: return self.get_full_url() @property - def host(self): + def host(self) -> str: return self.get_host() @property - def type(self): + def type(self) -> str: return self.get_type() @property - def unverifiable(self): + def unverifiable(self) -> bool: return self.is_unverifiable() @property - def origin_req_host(self): - return urlparse_cached(self.request).hostname + def origin_req_host(self) -> str: + return cast(str, urlparse_cached(self.request).hostname) - def has_header(self, name): + def has_header(self, name: str) -> bool: return name in self.request.headers - def get_header(self, name, default=None): + def get_header(self, name: str, default: Optional[str] = None) -> Optional[str]: value = self.request.headers.get(name, default) return to_unicode(value, errors="replace") if value is not None else None - def header_items(self): + def header_items(self) -> List[Tuple[str, List[str]]]: return [ ( to_unicode(k, errors="replace"), @@ -182,18 +207,18 @@ class WrappedRequest: for k, v in self.request.headers.items() ] - def add_unredirected_header(self, name, value): + def add_unredirected_header(self, name: str, value: str) -> None: self.request.headers.appendlist(name, value) class WrappedResponse: - def __init__(self, response): + def __init__(self, response: Response): self.response = response - def info(self): + def info(self) -> Self: return self - def get_all(self, name, default=None): + def get_all(self, name: str, default: Any = None) -> List[str]: return [ to_unicode(v, errors="replace") for v in self.response.headers.getlist(name) ] diff --git a/scrapy/http/response/text.py b/scrapy/http/response/text.py index 2816610fb..522ffc0d5 100644 --- a/scrapy/http/response/text.py +++ b/scrapy/http/response/text.py @@ -159,12 +159,12 @@ class TextResponse(Response): def jmespath(self, query: str, **kwargs: Any) -> SelectorList: from scrapy.selector import SelectorList - if not hasattr(self.selector, "jmespath"): # type: ignore[attr-defined] + if not hasattr(self.selector, "jmespath"): raise AttributeError( "Please install parsel >= 1.8.1 to get jmespath support" ) - return cast(SelectorList, self.selector.jmespath(query, **kwargs)) # type: ignore[attr-defined] + return cast(SelectorList, self.selector.jmespath(query, **kwargs)) def xpath(self, query: str, **kwargs: Any) -> SelectorList: from scrapy.selector import SelectorList