Typing for scrapy/core/downloader (#5896)

This commit is contained in:
Andrey Rakhmatullin 2023-04-17 10:37:52 +04:00 committed by GitHub
parent e1f66620ec
commit 02f3e8d413
No known key found for this signature in database
GPG Key ID: 4AEE18F83AFDEB23
8 changed files with 211 additions and 125 deletions

View File

@ -2,45 +2,52 @@ import random
from collections import deque
from datetime import datetime
from time import time
from typing import TYPE_CHECKING, Any, Deque, Dict, Set, Tuple, cast
from twisted.internet import defer, task
from twisted.internet import task
from twisted.internet.defer import Deferred
from scrapy import Request, Spider, signals
from scrapy.core.downloader.handlers import DownloadHandlers
from scrapy.core.downloader.middleware import DownloaderMiddlewareManager
from scrapy.http import Response
from scrapy.resolver import dnscache
from scrapy.settings import BaseSettings
from scrapy.signalmanager import SignalManager
from scrapy.utils.defer import mustbe_deferred
from scrapy.utils.httpobj import urlparse_cached
if TYPE_CHECKING:
from scrapy.crawler import Crawler
class Slot:
"""Downloader slot"""
def __init__(self, concurrency, delay, randomize_delay):
self.concurrency = concurrency
self.delay = delay
self.randomize_delay = randomize_delay
def __init__(self, concurrency: int, delay: float, randomize_delay: bool):
self.concurrency: int = concurrency
self.delay: float = delay
self.randomize_delay: bool = randomize_delay
self.active = set()
self.queue = deque()
self.transferring = set()
self.lastseen = 0
self.active: Set[Request] = set()
self.queue: Deque[Tuple[Request, Deferred]] = deque()
self.transferring: Set[Request] = set()
self.lastseen: float = 0
self.latercall = None
def free_transfer_slots(self):
def free_transfer_slots(self) -> int:
return self.concurrency - len(self.transferring)
def download_delay(self):
def download_delay(self) -> float:
if self.randomize_delay:
return random.uniform(0.5 * self.delay, 1.5 * self.delay)
return self.delay
def close(self):
def close(self) -> None:
if self.latercall and self.latercall.active():
self.latercall.cancel()
def __repr__(self):
def __repr__(self) -> str:
cls_name = self.__class__.__name__
return (
f"{cls_name}(concurrency={self.concurrency!r}, "
@ -48,7 +55,7 @@ class Slot:
f"randomize_delay={self.randomize_delay!r})"
)
def __str__(self):
def __str__(self) -> str:
return (
f"<downloader.Slot concurrency={self.concurrency!r} "
f"delay={self.delay:.2f} randomize_delay={self.randomize_delay!r} "
@ -58,8 +65,10 @@ class Slot:
)
def _get_concurrency_delay(concurrency, spider, settings):
delay = settings.getfloat("DOWNLOAD_DELAY")
def _get_concurrency_delay(
concurrency: int, spider: Spider, settings: BaseSettings
) -> Tuple[int, float]:
delay: float = settings.getfloat("DOWNLOAD_DELAY")
if hasattr(spider, "download_delay"):
delay = spider.download_delay
@ -72,23 +81,29 @@ def _get_concurrency_delay(concurrency, spider, settings):
class Downloader:
DOWNLOAD_SLOT = "download_slot"
def __init__(self, crawler):
self.settings = crawler.settings
self.signals = crawler.signals
self.slots = {}
self.active = set()
self.handlers = DownloadHandlers(crawler)
self.total_concurrency = self.settings.getint("CONCURRENT_REQUESTS")
self.domain_concurrency = self.settings.getint("CONCURRENT_REQUESTS_PER_DOMAIN")
self.ip_concurrency = self.settings.getint("CONCURRENT_REQUESTS_PER_IP")
self.randomize_delay = self.settings.getbool("RANDOMIZE_DOWNLOAD_DELAY")
self.middleware = DownloaderMiddlewareManager.from_crawler(crawler)
self._slot_gc_loop = task.LoopingCall(self._slot_gc)
def __init__(self, crawler: "Crawler"):
self.settings: BaseSettings = crawler.settings
self.signals: SignalManager = crawler.signals
self.slots: Dict[str, Slot] = {}
self.active: Set[Request] = set()
self.handlers: DownloadHandlers = DownloadHandlers(crawler)
self.total_concurrency: int = self.settings.getint("CONCURRENT_REQUESTS")
self.domain_concurrency: int = self.settings.getint(
"CONCURRENT_REQUESTS_PER_DOMAIN"
)
self.ip_concurrency: int = self.settings.getint("CONCURRENT_REQUESTS_PER_IP")
self.randomize_delay: bool = self.settings.getbool("RANDOMIZE_DOWNLOAD_DELAY")
self.middleware: DownloaderMiddlewareManager = (
DownloaderMiddlewareManager.from_crawler(crawler)
)
self._slot_gc_loop: task.LoopingCall = task.LoopingCall(self._slot_gc)
self._slot_gc_loop.start(60)
self.per_slot_settings = self.settings.getdict("DOWNLOAD_SLOTS", {})
self.per_slot_settings: Dict[str, Dict[str, Any]] = self.settings.getdict(
"DOWNLOAD_SLOTS", {}
)
def fetch(self, request: Request, spider: Spider) -> Deferred:
def _deactivate(response):
def _deactivate(response: Response) -> Response:
self.active.remove(request)
return response
@ -99,7 +114,7 @@ class Downloader:
def needs_backout(self) -> bool:
return len(self.active) >= self.total_concurrency
def _get_slot(self, request, spider):
def _get_slot(self, request: Request, spider: Spider) -> Tuple[str, Slot]:
key = self._get_slot_key(request, spider)
if key not in self.slots:
slot_settings = self.per_slot_settings.get(key, {})
@ -117,9 +132,9 @@ class Downloader:
return key, self.slots[key]
def _get_slot_key(self, request, spider):
def _get_slot_key(self, request: Request, spider: Spider) -> str:
if self.DOWNLOAD_SLOT in request.meta:
return request.meta[self.DOWNLOAD_SLOT]
return cast(str, request.meta[self.DOWNLOAD_SLOT])
key = urlparse_cached(request).hostname or ""
if self.ip_concurrency:
@ -127,11 +142,11 @@ class Downloader:
return key
def _enqueue_request(self, request, spider):
def _enqueue_request(self, request: Request, spider: Spider) -> Deferred:
key, slot = self._get_slot(request, spider)
request.meta[self.DOWNLOAD_SLOT] = key
def _deactivate(response):
def _deactivate(response: Response) -> Response:
slot.active.remove(request)
return response
@ -139,12 +154,12 @@ class Downloader:
self.signals.send_catch_log(
signal=signals.request_reached_downloader, request=request, spider=spider
)
deferred = defer.Deferred().addBoth(_deactivate)
deferred = Deferred().addBoth(_deactivate)
slot.queue.append((request, deferred))
self._process_queue(spider, slot)
return deferred
def _process_queue(self, spider, slot):
def _process_queue(self, spider: Spider, slot: Slot) -> None:
from twisted.internet import reactor
if slot.latercall and slot.latercall.active():
@ -172,7 +187,7 @@ class Downloader:
self._process_queue(spider, slot)
break
def _download(self, slot, request, spider):
def _download(self, slot: Slot, request: Request, spider: Spider) -> Deferred:
# The order is very important for the following deferreds. Do not change!
# 1. Create the download deferred
@ -180,7 +195,7 @@ class Downloader:
# 2. Notify response_downloaded listeners about the recent download
# before querying queue for next request
def _downloaded(response):
def _downloaded(response: Response) -> Response:
self.signals.send_catch_log(
signal=signals.response_downloaded,
response=response,
@ -197,7 +212,7 @@ class Downloader:
# middleware itself)
slot.transferring.add(request)
def finish_transferring(_):
def finish_transferring(_: Any) -> Any:
slot.transferring.remove(request)
self._process_queue(spider, slot)
self.signals.send_catch_log(

View File

@ -1,4 +1,5 @@
import warnings
from typing import TYPE_CHECKING, Any, List, Optional
from OpenSSL import SSL
from twisted.internet._sslverify import _setAcceptableProtocols
@ -18,8 +19,12 @@ from scrapy.core.downloader.tls import (
ScrapyClientTLSOptions,
openssl_methods,
)
from scrapy.settings import BaseSettings
from scrapy.utils.misc import create_instance, load_object
if TYPE_CHECKING:
from twisted.internet._sslverify import ClientTLSOptions
@implementer(IPolicyForHTTPS)
class ScrapyClientContextFactory(BrowserLikePolicyForHTTPS):
@ -35,25 +40,34 @@ class ScrapyClientContextFactory(BrowserLikePolicyForHTTPS):
def __init__(
self,
method=SSL.SSLv23_METHOD,
tls_verbose_logging=False,
tls_ciphers=None,
*args,
**kwargs,
method: int = SSL.SSLv23_METHOD,
tls_verbose_logging: bool = False,
tls_ciphers: Optional[str] = None,
*args: Any,
**kwargs: Any,
):
super().__init__(*args, **kwargs)
self._ssl_method = method
self.tls_verbose_logging = tls_verbose_logging
self._ssl_method: int = method
self.tls_verbose_logging: bool = tls_verbose_logging
self.tls_ciphers: AcceptableCiphers
if tls_ciphers:
self.tls_ciphers = AcceptableCiphers.fromOpenSSLCipherString(tls_ciphers)
else:
self.tls_ciphers = DEFAULT_CIPHERS
@classmethod
def from_settings(cls, settings, method=SSL.SSLv23_METHOD, *args, **kwargs):
tls_verbose_logging = settings.getbool("DOWNLOADER_CLIENT_TLS_VERBOSE_LOGGING")
tls_ciphers = settings["DOWNLOADER_CLIENT_TLS_CIPHERS"]
return cls(
def from_settings(
cls,
settings: BaseSettings,
method: int = SSL.SSLv23_METHOD,
*args: Any,
**kwargs: Any,
):
tls_verbose_logging: bool = settings.getbool(
"DOWNLOADER_CLIENT_TLS_VERBOSE_LOGGING"
)
tls_ciphers: Optional[str] = settings["DOWNLOADER_CLIENT_TLS_CIPHERS"]
return cls( # type: ignore[misc]
method=method,
tls_verbose_logging=tls_verbose_logging,
tls_ciphers=tls_ciphers,
@ -61,7 +75,7 @@ class ScrapyClientContextFactory(BrowserLikePolicyForHTTPS):
**kwargs,
)
def getCertificateOptions(self):
def getCertificateOptions(self) -> CertificateOptions:
# setting verify=True will require you to provide CAs
# to verify against; in other words: it's not that simple
@ -82,12 +96,12 @@ class ScrapyClientContextFactory(BrowserLikePolicyForHTTPS):
# kept for old-style HTTP/1.0 downloader context twisted calls,
# e.g. connectSSL()
def getContext(self, hostname=None, port=None):
def getContext(self, hostname: Any = None, port: Any = None) -> SSL.Context:
ctx = self.getCertificateOptions().getContext()
ctx.set_options(0x4) # OP_LEGACY_SERVER_CONNECT
return ctx
def creatorForNetloc(self, hostname, port):
def creatorForNetloc(self, hostname: bytes, port: int) -> "ClientTLSOptions":
return ScrapyClientTLSOptions(
hostname.decode("ascii"),
self.getContext(),
@ -114,7 +128,7 @@ class BrowserLikeContextFactory(ScrapyClientContextFactory):
``SSLv23_METHOD``) which allows TLS protocol negotiation.
"""
def creatorForNetloc(self, hostname, port):
def creatorForNetloc(self, hostname: bytes, port: int) -> "ClientTLSOptions":
# trustRoot set to platformTrust() will use the platform's root CAs.
#
# This means that a website like https://www.cacert.org will be rejected
@ -133,13 +147,15 @@ class AcceptableProtocolsContextFactory:
negotiation.
"""
def __init__(self, context_factory, acceptable_protocols):
def __init__(self, context_factory: Any, acceptable_protocols: List[bytes]):
verifyObject(IPolicyForHTTPS, context_factory)
self._wrapped_context_factory = context_factory
self._acceptable_protocols = acceptable_protocols
self._wrapped_context_factory: Any = context_factory
self._acceptable_protocols: List[bytes] = acceptable_protocols
def creatorForNetloc(self, hostname, port):
options = self._wrapped_context_factory.creatorForNetloc(hostname, port)
def creatorForNetloc(self, hostname: bytes, port: int) -> "ClientTLSOptions":
options: "ClientTLSOptions" = self._wrapped_context_factory.creatorForNetloc(
hostname, port
)
_setAcceptableProtocols(options._ctx, self._acceptable_protocols)
return options

View File

@ -1,25 +1,32 @@
"""Download handlers for different schemes"""
import logging
from typing import TYPE_CHECKING, Any, Callable, Dict, Generator, Union, cast
from twisted.internet import defer
from twisted.internet.defer import Deferred
from scrapy import signals
from scrapy import Request, Spider, signals
from scrapy.exceptions import NotConfigured, NotSupported
from scrapy.utils.httpobj import urlparse_cached
from scrapy.utils.misc import create_instance, load_object
from scrapy.utils.python import without_none_values
if TYPE_CHECKING:
from scrapy.crawler import Crawler
logger = logging.getLogger(__name__)
class DownloadHandlers:
def __init__(self, crawler):
self._crawler = crawler
self._schemes = {} # stores acceptable schemes on instancing
self._handlers = {} # stores instanced handlers for schemes
self._notconfigured = {} # remembers failed handlers
handlers = without_none_values(
def __init__(self, crawler: "Crawler"):
self._crawler: "Crawler" = crawler
self._schemes: Dict[
str, Union[str, Callable]
] = {} # stores acceptable schemes on instancing
self._handlers: Dict[str, Any] = {} # stores instanced handlers for schemes
self._notconfigured: Dict[str, str] = {} # remembers failed handlers
handlers: Dict[str, Union[str, Callable]] = without_none_values(
crawler.settings.getwithbase("DOWNLOAD_HANDLERS")
)
for scheme, clspath in handlers.items():
@ -28,7 +35,7 @@ class DownloadHandlers:
crawler.signals.connect(self._close, signals.engine_stopped)
def _get_handler(self, scheme):
def _get_handler(self, scheme: str) -> Any:
"""Lazy-load the downloadhandler for a scheme
only on the first request for that scheme.
"""
@ -42,7 +49,7 @@ class DownloadHandlers:
return self._load_handler(scheme)
def _load_handler(self, scheme, skip_lazy=False):
def _load_handler(self, scheme: str, skip_lazy: bool = False) -> Any:
path = self._schemes[scheme]
try:
dhcls = load_object(path)
@ -69,17 +76,17 @@ class DownloadHandlers:
self._handlers[scheme] = dh
return dh
def download_request(self, request, spider):
def download_request(self, request: Request, spider: Spider) -> Deferred:
scheme = urlparse_cached(request).scheme
handler = self._get_handler(scheme)
if not handler:
raise NotSupported(
f"Unsupported URL scheme '{scheme}': {self._notconfigured[scheme]}"
)
return handler.download_request(request, spider)
return cast(Deferred, handler.download_request(request, spider))
@defer.inlineCallbacks
def _close(self, *_a, **_kw):
def _close(self, *_a: Any, **_kw: Any) -> Generator[Deferred, Any, None]:
for dh in self._handlers.values():
if hasattr(dh, "close"):
yield dh.close()

View File

@ -3,15 +3,16 @@ Downloader Middleware manager
See documentation in docs/topics/downloader-middleware.rst
"""
from typing import Callable, Union, cast
from typing import Any, Callable, Generator, List, Union, cast
from twisted.internet import defer
from twisted.internet.defer import Deferred, inlineCallbacks
from twisted.python.failure import Failure
from scrapy import Spider
from scrapy.exceptions import _InvalidOutput
from scrapy.http import Request, Response
from scrapy.middleware import MiddlewareManager
from scrapy.settings import BaseSettings
from scrapy.utils.conf import build_component_list
from scrapy.utils.defer import deferred_from_coro, mustbe_deferred
@ -20,10 +21,10 @@ class DownloaderMiddlewareManager(MiddlewareManager):
component_name = "downloader middleware"
@classmethod
def _get_mwlist_from_settings(cls, settings):
def _get_mwlist_from_settings(cls, settings: BaseSettings) -> List[Any]:
return build_component_list(settings.getwithbase("DOWNLOADER_MIDDLEWARES"))
def _add_middleware(self, mw):
def _add_middleware(self, mw: Any) -> None:
if hasattr(mw, "process_request"):
self.methods["process_request"].append(mw.process_request)
if hasattr(mw, "process_response"):
@ -31,9 +32,11 @@ class DownloaderMiddlewareManager(MiddlewareManager):
if hasattr(mw, "process_exception"):
self.methods["process_exception"].appendleft(mw.process_exception)
def download(self, download_func: Callable, request: Request, spider: Spider):
@defer.inlineCallbacks
def process_request(request: Request):
def download(
self, download_func: Callable, request: Request, spider: Spider
) -> Deferred:
@inlineCallbacks
def process_request(request: Request) -> Generator[Deferred, Any, Any]:
for method in self.methods["process_request"]:
method = cast(Callable, method)
response = yield deferred_from_coro(
@ -50,8 +53,10 @@ class DownloaderMiddlewareManager(MiddlewareManager):
return response
return (yield download_func(request=request, spider=spider))
@defer.inlineCallbacks
def process_response(response: Union[Response, Request]):
@inlineCallbacks
def process_response(
response: Union[Response, Request]
) -> Generator[Deferred, Any, Union[Response, Request]]:
if response is None:
raise TypeError("Received None in process_response")
elif isinstance(response, Request):
@ -71,8 +76,10 @@ class DownloaderMiddlewareManager(MiddlewareManager):
return response
return response
@defer.inlineCallbacks
def process_exception(failure: Failure):
@inlineCallbacks
def process_exception(
failure: Failure,
) -> Generator[Deferred, Any, Union[Failure, Response, Request]]:
exception = failure.value
for method in self.methods["process_exception"]:
method = cast(Callable, method)

View File

@ -1,4 +1,5 @@
import logging
from typing import Any, Dict
from OpenSSL import SSL
from service_identity.exceptions import CertificateError
@ -20,7 +21,7 @@ METHOD_TLSv11 = "TLSv1.1"
METHOD_TLSv12 = "TLSv1.2"
openssl_methods = {
openssl_methods: Dict[str, int] = {
METHOD_TLS: SSL.SSLv23_METHOD, # protocol negotiation (recommended)
METHOD_TLSv10: SSL.TLSv1_METHOD, # TLS 1.0 only
METHOD_TLSv11: SSL.TLSv1_1_METHOD, # TLS 1.1 only
@ -39,11 +40,13 @@ class ScrapyClientTLSOptions(ClientTLSOptions):
logging warnings. Also, HTTPS connection parameters logging is added.
"""
def __init__(self, hostname, ctx, verbose_logging=False):
def __init__(self, hostname: str, ctx: SSL.Context, verbose_logging: bool = False):
super().__init__(hostname, ctx)
self.verbose_logging = verbose_logging
self.verbose_logging: bool = verbose_logging
def _identityVerifyingInfoCallback(self, connection, where, ret):
def _identityVerifyingInfoCallback(
self, connection: SSL.Connection, where: int, ret: Any
) -> None:
if where & SSL.SSL_CB_HANDSHAKE_START:
connection.set_tlsext_host_name(self._hostnameBytes)
elif where & SSL.SSL_CB_HANDSHAKE_DONE:
@ -55,11 +58,12 @@ class ScrapyClientTLSOptions(ClientTLSOptions):
connection.get_cipher_name(),
)
server_cert = connection.get_peer_certificate()
logger.debug(
'SSL connection certificate: issuer "%s", subject "%s"',
x509name_to_string(server_cert.get_issuer()),
x509name_to_string(server_cert.get_subject()),
)
if server_cert:
logger.debug(
'SSL connection certificate: issuer "%s", subject "%s"',
x509name_to_string(server_cert.get_issuer()),
x509name_to_string(server_cert.get_subject()),
)
key_info = get_temp_key_info(connection._ssl)
if key_info:
logger.debug("SSL temp key: %s", key_info)
@ -82,4 +86,6 @@ class ScrapyClientTLSOptions(ClientTLSOptions):
)
DEFAULT_CIPHERS = AcceptableCiphers.fromOpenSSLCipherString("DEFAULT")
DEFAULT_CIPHERS: AcceptableCiphers = AcceptableCiphers.fromOpenSSLCipherString(
"DEFAULT"
)

View File

@ -1,22 +1,25 @@
import re
from time import time
from urllib.parse import urldefrag, urlparse, urlunparse
from typing import Optional, Tuple
from urllib.parse import ParseResult, urldefrag, urlparse, urlunparse
from twisted.internet import defer
from twisted.internet.protocol import ClientFactory
from twisted.web.http import HTTPClient
from scrapy import Request
from scrapy.http import Headers
from scrapy.responsetypes import responsetypes
from scrapy.utils.httpobj import urlparse_cached
from scrapy.utils.python import to_bytes, to_unicode
def _parsed_url_args(parsed):
def _parsed_url_args(parsed: ParseResult) -> Tuple[bytes, bytes, bytes, int, bytes]:
# Assume parsed is urlparse-d from Request.url,
# which was passed via safe_url_string and is ascii-only.
path = urlunparse(("", "", parsed.path or "/", parsed.params, parsed.query, ""))
path = to_bytes(path, encoding="ascii")
path_str = urlunparse(("", "", parsed.path or "/", parsed.params, parsed.query, ""))
path = to_bytes(path_str, encoding="ascii")
assert parsed.hostname is not None
host = to_bytes(parsed.hostname, encoding="ascii")
port = parsed.port
scheme = to_bytes(parsed.scheme, encoding="ascii")
@ -26,7 +29,7 @@ def _parsed_url_args(parsed):
return scheme, netloc, host, port, path
def _parse(url):
def _parse(url: str) -> Tuple[bytes, bytes, bytes, int, bytes]:
"""Return tuple of (scheme, netloc, host, port, path),
all in bytes except for port which is int.
Assume url is from Request.url, which was passed via safe_url_string
@ -132,17 +135,19 @@ class ScrapyHTTPClientFactory(ClientFactory):
self.scheme, _, self.host, self.port, _ = _parse(proxy)
self.path = self.url
def __init__(self, request, timeout=180):
self._url = urldefrag(request.url)[0]
def __init__(self, request: Request, timeout: float = 180):
self._url: str = urldefrag(request.url)[0]
# converting to bytes to comply to Twisted interface
self.url = to_bytes(self._url, encoding="ascii")
self.method = to_bytes(request.method, encoding="ascii")
self.body = request.body or None
self.headers = Headers(request.headers)
self.response_headers = None
self.timeout = request.meta.get("download_timeout") or timeout
self.start_time = time()
self.deferred = defer.Deferred().addCallback(self._build_response, request)
self.url: bytes = to_bytes(self._url, encoding="ascii")
self.method: bytes = to_bytes(request.method, encoding="ascii")
self.body: Optional[bytes] = request.body or None
self.headers: Headers = Headers(request.headers)
self.response_headers: Optional[Headers] = None
self.timeout: float = request.meta.get("download_timeout") or timeout
self.start_time: float = time()
self.deferred: defer.Deferred = defer.Deferred().addCallback(
self._build_response, request
)
# Fixes Twisted 11.1.0+ support as HTTPClientFactory is expected
# to have _disconnectedDeferred. See Twisted r32329.
@ -150,7 +155,7 @@ class ScrapyHTTPClientFactory(ClientFactory):
# needed to add the callback _waitForDisconnect.
# Specifically this avoids the AttributeError exception when
# clientConnectionFailed method is called.
self._disconnectedDeferred = defer.Deferred()
self._disconnectedDeferred: defer.Deferred = defer.Deferred()
self._set_connection_attributes(request)
@ -166,8 +171,8 @@ class ScrapyHTTPClientFactory(ClientFactory):
elif self.method == b"POST":
self.headers["Content-Length"] = 0
def __repr__(self):
return f"<{self.__class__.__name__}: {self.url}>"
def __repr__(self) -> str:
return f"<{self.__class__.__name__}: {self._url}>"
def _cancelTimeout(self, result, timeoutCall):
if timeoutCall.active():

View File

@ -8,7 +8,16 @@ import sys
import weakref
from functools import partial, wraps
from itertools import chain
from typing import Any, AsyncGenerator, AsyncIterable, Iterable, Union
from typing import (
Any,
AsyncGenerator,
AsyncIterable,
Iterable,
Mapping,
Optional,
Union,
overload,
)
from scrapy.utils.asyncgen import as_async_generator
@ -82,7 +91,9 @@ def unique(list_, key=lambda x: x):
return result
def to_unicode(text, encoding=None, errors="strict"):
def to_unicode(
text: Union[str, bytes], encoding: Optional[str] = None, errors: str = "strict"
) -> str:
"""Return the unicode representation of a bytes object ``text``. If
``text`` is already an unicode object, return it as-is."""
if isinstance(text, str):
@ -97,7 +108,9 @@ def to_unicode(text, encoding=None, errors="strict"):
return text.decode(encoding, errors)
def to_bytes(text, encoding=None, errors="strict"):
def to_bytes(
text: Union[str, bytes], encoding: Optional[str] = None, errors: str = "strict"
) -> bytes:
"""Return the binary representation of ``text``. If ``text``
is already a bytes object, return it as-is."""
if isinstance(text, bytes):
@ -160,11 +173,12 @@ def memoizemethod_noargs(method):
return new_method
_BINARYCHARS = {to_bytes(chr(i)) for i in range(32)} - {b"\0", b"\t", b"\n", b"\r"}
_BINARYCHARS |= {ord(ch) for ch in _BINARYCHARS}
_BINARYCHARS = {
i for i in range(32) if to_bytes(chr(i)) not in {b"\0", b"\t", b"\n", b"\r"}
}
def binary_is_text(data):
def binary_is_text(data: bytes) -> bool:
"""Returns ``True`` if the given ``data`` argument (a ``bytes`` object)
does not contain unprintable control characters.
"""
@ -258,6 +272,16 @@ def equal_attributes(obj1, obj2, attributes):
return True
@overload
def without_none_values(iterable: Mapping) -> dict:
...
@overload
def without_none_values(iterable: Iterable) -> Iterable:
...
def without_none_values(iterable):
"""Return a copy of ``iterable`` with all ``None`` entries removed.

View File

@ -1,24 +1,28 @@
from typing import Any, Optional, cast
import OpenSSL._util as pyOpenSSLutil
import OpenSSL.SSL
import OpenSSL.version
from OpenSSL.crypto import X509Name
from scrapy.utils.python import to_unicode
def ffi_buf_to_string(buf):
def ffi_buf_to_string(buf: Any) -> str:
return to_unicode(pyOpenSSLutil.ffi.string(buf))
def x509name_to_string(x509name):
def x509name_to_string(x509name: X509Name) -> str:
# from OpenSSL.crypto.X509Name.__repr__
result_buffer = pyOpenSSLutil.ffi.new("char[]", 512)
result_buffer: Any = pyOpenSSLutil.ffi.new("char[]", 512)
pyOpenSSLutil.lib.X509_NAME_oneline(
x509name._name, result_buffer, len(result_buffer)
x509name._name, result_buffer, len(result_buffer) # type: ignore[attr-defined]
)
return ffi_buf_to_string(result_buffer)
def get_temp_key_info(ssl_object):
def get_temp_key_info(ssl_object: Any) -> Optional[str]:
# adapted from OpenSSL apps/s_cb.c::ssl_print_tmp_key()
if not hasattr(pyOpenSSLutil.lib, "SSL_get_server_tmp_key"):
# removed in cryptography 40.0.0
@ -53,8 +57,10 @@ def get_temp_key_info(ssl_object):
return ", ".join(key_info)
def get_openssl_version():
system_openssl = OpenSSL.SSL.SSLeay_version(OpenSSL.SSL.SSLEAY_VERSION).decode(
"ascii", errors="replace"
def get_openssl_version() -> str:
# https://github.com/python/typeshed/issues/10024
system_openssl_bytes = cast(
bytes, OpenSSL.SSL.SSLeay_version(OpenSSL.SSL.SSLEAY_VERSION)
)
system_openssl = system_openssl_bytes.decode("ascii", errors="replace")
return f"{OpenSSL.version.__version__} ({system_openssl})"