Merge pull request #6341 from wRAR/typing-downloader

More typing for scrapy/core/downloader
This commit is contained in:
Andrey Rakhmatullin 2024-05-07 12:24:25 +04:00 committed by GitHub
commit c9bac7a657
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
7 changed files with 240 additions and 162 deletions

View File

@ -21,6 +21,7 @@ from scrapy.core.downloader.tls import (
ScrapyClientTLSOptions,
openssl_methods,
)
from scrapy.crawler import Crawler
from scrapy.settings import BaseSettings
from scrapy.utils.misc import build_from_crawler, load_object
@ -102,7 +103,7 @@ class ScrapyClientContextFactory(BrowserLikePolicyForHTTPS):
# kept for old-style HTTP/1.0 downloader context twisted calls,
# e.g. connectSSL()
def getContext(self, hostname: Any = None, port: Any = None) -> SSL.Context:
ctx = self.getCertificateOptions().getContext()
ctx: SSL.Context = self.getCertificateOptions().getContext()
ctx.set_options(0x4) # OP_LEGACY_SERVER_CONNECT
return ctx
@ -165,7 +166,9 @@ class AcceptableProtocolsContextFactory:
return options
def load_context_factory_from_settings(settings, crawler):
def load_context_factory_from_settings(
settings: BaseSettings, crawler: Crawler
) -> IPolicyForHTTPS:
ssl_method = openssl_methods[settings.get("DOWNLOADER_CLIENT_TLS_METHOD")]
context_factory_cls = load_object(settings["DOWNLOADER_CLIENTCONTEXTFACTORY"])
# try method-aware context factory

View File

@ -2,6 +2,8 @@ from pathlib import Path
from w3lib.url import file_uri_to_path
from scrapy import Request, Spider
from scrapy.http import Response
from scrapy.responsetypes import responsetypes
from scrapy.utils.decorators import defers
@ -10,7 +12,7 @@ class FileDownloadHandler:
lazy = False
@defers
def download_request(self, request, spider):
def download_request(self, request: Request, spider: Spider) -> Response:
filepath = file_uri_to_path(request.url)
body = Path(filepath).read_bytes()
respcls = responsetypes.from_args(filename=filepath, body=body)

View File

@ -32,15 +32,19 @@ from __future__ import annotations
import re
from io import BytesIO
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any, BinaryIO, Dict, Optional
from urllib.parse import unquote
from twisted.internet.defer import Deferred
from twisted.internet.protocol import ClientCreator, Protocol
from twisted.protocols.ftp import CommandFailed, FTPClient
from twisted.python.failure import Failure
from scrapy import Request, Spider
from scrapy.crawler import Crawler
from scrapy.http import Response
from scrapy.responsetypes import responsetypes
from scrapy.settings import BaseSettings
from scrapy.utils.httpobj import urlparse_cached
from scrapy.utils.python import to_bytes
@ -50,20 +54,20 @@ if TYPE_CHECKING:
class ReceivedDataProtocol(Protocol):
def __init__(self, filename=None):
self.__filename = filename
self.body = open(filename, "wb") if filename else BytesIO()
self.size = 0
def __init__(self, filename: Optional[str] = None):
self.__filename: Optional[str] = filename
self.body: BinaryIO = open(filename, "wb") if filename else BytesIO()
self.size: int = 0
def dataReceived(self, data):
def dataReceived(self, data: bytes) -> None:
self.body.write(data)
self.size += len(data)
@property
def filename(self):
def filename(self) -> Optional[str]:
return self.__filename
def close(self):
def close(self) -> None:
self.body.close() if self.filename else self.body.seek(0)
@ -73,12 +77,12 @@ _CODE_RE = re.compile(r"\d+")
class FTPDownloadHandler:
lazy = False
CODE_MAPPING = {
CODE_MAPPING: Dict[str, int] = {
"550": 404,
"default": 503,
}
def __init__(self, settings):
def __init__(self, settings: BaseSettings):
self.default_user = settings["FTP_USER"]
self.default_password = settings["FTP_PASSWORD"]
self.passive_mode = settings["FTP_PASSIVE_MODE"]
@ -87,7 +91,7 @@ class FTPDownloadHandler:
def from_crawler(cls, crawler: Crawler) -> Self:
return cls(crawler.settings)
def download_request(self, request, spider):
def download_request(self, request: Request, spider: Spider) -> Deferred:
from twisted.internet import reactor
parsed_url = urlparse_cached(request)
@ -99,10 +103,10 @@ class FTPDownloadHandler:
creator = ClientCreator(
reactor, FTPClient, user, password, passive=passive_mode
)
dfd = creator.connectTCP(parsed_url.hostname, parsed_url.port or 21)
dfd: Deferred = creator.connectTCP(parsed_url.hostname, parsed_url.port or 21)
return dfd.addCallback(self.gotClient, request, unquote(parsed_url.path))
def gotClient(self, client, request, filepath):
def gotClient(self, client: FTPClient, request: Request, filepath: str) -> Deferred:
self.client = client
protocol = ReceivedDataProtocol(request.meta.get("ftp_local_filename"))
return client.retrieveFile(filepath, protocol).addCallbacks(
@ -112,15 +116,18 @@ class FTPDownloadHandler:
errbackArgs=(request,),
)
def _build_response(self, result, request, protocol):
def _build_response(
self, result: Any, request: Request, protocol: ReceivedDataProtocol
) -> Response:
self.result = result
protocol.close()
headers = {"local filename": protocol.filename or "", "size": protocol.size}
body = to_bytes(protocol.filename or protocol.body.read())
respcls = responsetypes.from_args(url=request.url, body=body)
return respcls(url=request.url, status=200, body=body, headers=headers)
# hints for Headers-related types may need to be fixed to not use AnyStr
return respcls(url=request.url, status=200, body=body, headers=headers) # type: ignore[arg-type]
def _failed(self, result, request):
def _failed(self, result: Failure, request: Request) -> Response:
message = result.getErrorMessage()
if result.type == CommandFailed:
m = _CODE_RE.search(message)
@ -130,4 +137,5 @@ class FTPDownloadHandler:
return Response(
url=request.url, status=httpcode, body=to_bytes(message)
)
assert result.type
raise result.type(result.value)

View File

@ -3,8 +3,13 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Type
from twisted.internet.defer import Deferred
from scrapy import Request, Spider
from scrapy.crawler import Crawler
from scrapy.settings import BaseSettings
from scrapy.utils.misc import build_from_crawler, load_object
from scrapy.utils.python import to_unicode
@ -12,29 +17,34 @@ if TYPE_CHECKING:
# typing.Self requires Python 3.11
from typing_extensions import Self
from scrapy.core.downloader.contextfactory import ScrapyClientContextFactory
from scrapy.core.downloader.webclient import ScrapyHTTPClientFactory
class HTTP10DownloadHandler:
lazy = False
def __init__(self, settings, crawler=None):
self.HTTPClientFactory = load_object(settings["DOWNLOADER_HTTPCLIENTFACTORY"])
self.ClientContextFactory = load_object(
def __init__(self, settings: BaseSettings, crawler: Crawler):
self.HTTPClientFactory: Type[ScrapyHTTPClientFactory] = load_object(
settings["DOWNLOADER_HTTPCLIENTFACTORY"]
)
self.ClientContextFactory: Type[ScrapyClientContextFactory] = load_object(
settings["DOWNLOADER_CLIENTCONTEXTFACTORY"]
)
self._settings = settings
self._crawler = crawler
self._settings: BaseSettings = settings
self._crawler: Crawler = crawler
@classmethod
def from_crawler(cls, crawler) -> Self:
def from_crawler(cls, crawler: Crawler) -> Self:
return cls(crawler.settings, crawler)
def download_request(self, request, spider):
def download_request(self, request: Request, spider: Spider) -> Deferred:
"""Return a deferred for the HTTP download"""
factory = self.HTTPClientFactory(request)
self._connect(factory)
return factory.deferred
def _connect(self, factory):
def _connect(self, factory: ScrapyHTTPClientFactory) -> Deferred:
from twisted.internet import reactor
host, port = to_unicode(factory.host), factory.port

View File

@ -8,31 +8,33 @@ import re
from contextlib import suppress
from io import BytesIO
from time import time
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
from urllib.parse import urldefrag, urlunparse
from twisted.internet import defer, protocol, ssl
from twisted.internet import ssl
from twisted.internet.base import ReactorBase
from twisted.internet.defer import CancelledError, Deferred, succeed
from twisted.internet.endpoints import TCP4ClientEndpoint
from twisted.internet.error import TimeoutError
from twisted.internet.interfaces import IConsumer
from twisted.internet.protocol import Factory, Protocol, connectionDone
from twisted.python.failure import Failure
from twisted.web.client import (
URI,
Agent,
HTTPConnectionPool,
ResponseDone,
ResponseFailed,
)
from twisted.web.client import URI, Agent, HTTPConnectionPool
from twisted.web.client import Response as TxResponse
from twisted.web.client import ResponseDone, ResponseFailed
from twisted.web.http import PotentialDataLoss, _DataLoss
from twisted.web.http_headers import Headers as TxHeaders
from twisted.web.iweb import UNKNOWN_LENGTH, IBodyProducer
from twisted.web.iweb import UNKNOWN_LENGTH, IBodyProducer, IPolicyForHTTPS
from zope.interface import implementer
from scrapy import signals
from scrapy import Request, Spider, signals
from scrapy.core.downloader.contextfactory import load_context_factory_from_settings
from scrapy.core.downloader.webclient import _parse
from scrapy.crawler import Crawler
from scrapy.exceptions import StopDownload
from scrapy.http import Headers
from scrapy.http import Headers, Response
from scrapy.responsetypes import responsetypes
from scrapy.settings import BaseSettings
from scrapy.utils.python import to_bytes, to_unicode
if TYPE_CHECKING:
@ -46,28 +48,30 @@ logger = logging.getLogger(__name__)
class HTTP11DownloadHandler:
lazy = False
def __init__(self, settings, crawler=None):
def __init__(self, settings: BaseSettings, crawler: Crawler):
self._crawler = crawler
from twisted.internet import reactor
self._pool = HTTPConnectionPool(reactor, persistent=True)
self._pool: HTTPConnectionPool = HTTPConnectionPool(reactor, persistent=True)
self._pool.maxPersistentPerHost = settings.getint(
"CONCURRENT_REQUESTS_PER_DOMAIN"
)
self._pool._factory.noisy = False
self._contextFactory = load_context_factory_from_settings(settings, crawler)
self._default_maxsize = settings.getint("DOWNLOAD_MAXSIZE")
self._default_warnsize = settings.getint("DOWNLOAD_WARNSIZE")
self._fail_on_dataloss = settings.getbool("DOWNLOAD_FAIL_ON_DATALOSS")
self._disconnect_timeout = 1
self._contextFactory: IPolicyForHTTPS = load_context_factory_from_settings(
settings, crawler
)
self._default_maxsize: int = settings.getint("DOWNLOAD_MAXSIZE")
self._default_warnsize: int = settings.getint("DOWNLOAD_WARNSIZE")
self._fail_on_dataloss: bool = settings.getbool("DOWNLOAD_FAIL_ON_DATALOSS")
self._disconnect_timeout: int = 1
@classmethod
def from_crawler(cls, crawler) -> Self:
def from_crawler(cls, crawler: Crawler) -> Self:
return cls(crawler.settings, crawler)
def download_request(self, request, spider):
def download_request(self, request: Request, spider: Spider) -> Deferred:
"""Return a deferred for the HTTP download"""
agent = ScrapyAgent(
contextFactory=self._contextFactory,
@ -79,10 +83,10 @@ class HTTP11DownloadHandler:
)
return agent.download_request(request)
def close(self):
def close(self) -> Deferred:
from twisted.internet import reactor
d = self._pool.closeCachedConnections()
d: Deferred = self._pool.closeCachedConnections()
# closeCachedConnections will hang on network or server issues, so
# we'll manually timeout the deferred.
#
@ -93,7 +97,7 @@ class HTTP11DownloadHandler:
# issue a callback after `_disconnect_timeout` seconds.
delayed_call = reactor.callLater(self._disconnect_timeout, d.callback, [])
def cancel_delayed_call(result):
def cancel_delayed_call(result: Any) -> Any:
if delayed_call.active():
delayed_call.cancel()
return result
@ -123,39 +127,41 @@ class TunnelingTCP4ClientEndpoint(TCP4ClientEndpoint):
def __init__(
self,
reactor,
host,
port,
proxyConf,
contextFactory,
timeout=30,
bindAddress=None,
reactor: ReactorBase,
host: str,
port: int,
proxyConf: Tuple[str, int, Optional[bytes]],
contextFactory: IPolicyForHTTPS,
timeout: float = 30,
bindAddress: Optional[Tuple[str, int]] = None,
):
proxyHost, proxyPort, self._proxyAuthHeader = proxyConf
super().__init__(reactor, proxyHost, proxyPort, timeout, bindAddress)
self._tunnelReadyDeferred = defer.Deferred()
self._tunneledHost = host
self._tunneledPort = port
self._contextFactory = contextFactory
self._connectBuffer = bytearray()
self._tunnelReadyDeferred: Deferred = Deferred()
self._tunneledHost: str = host
self._tunneledPort: int = port
self._contextFactory: IPolicyForHTTPS = contextFactory
self._connectBuffer: bytearray = bytearray()
def requestTunnel(self, protocol):
def requestTunnel(self, protocol: Protocol) -> Protocol:
"""Asks the proxy to open a tunnel."""
assert protocol.transport
tunnelReq = tunnel_request_data(
self._tunneledHost, self._tunneledPort, self._proxyAuthHeader
)
protocol.transport.write(tunnelReq)
self._protocolDataReceived = protocol.dataReceived
protocol.dataReceived = self.processProxyResponse
protocol.dataReceived = self.processProxyResponse # type: ignore[method-assign]
self._protocol = protocol
return protocol
def processProxyResponse(self, rcvd_bytes):
def processProxyResponse(self, data: bytes) -> None:
"""Processes the response from the proxy. If the tunnel is successfully
created, notifies the client that we are ready to send requests. If not
raises a TunnelError.
"""
self._connectBuffer += rcvd_bytes
assert self._protocol.transport
self._connectBuffer += data
# make sure that enough (all) bytes are consumed
# and that we've got all HTTP headers (ending with a blank line)
# from the proxy so that we don't send those bytes to the TLS layer
@ -163,23 +169,24 @@ class TunnelingTCP4ClientEndpoint(TCP4ClientEndpoint):
# see https://github.com/scrapy/scrapy/issues/2491
if b"\r\n\r\n" not in self._connectBuffer:
return
self._protocol.dataReceived = self._protocolDataReceived
self._protocol.dataReceived = self._protocolDataReceived # type: ignore[method-assign]
respm = TunnelingTCP4ClientEndpoint._responseMatcher.match(self._connectBuffer)
if respm and int(respm.group("status")) == 200:
# set proper Server Name Indication extension
sslOptions = self._contextFactory.creatorForNetloc(
sslOptions = self._contextFactory.creatorForNetloc( # type: ignore[call-arg,misc]
self._tunneledHost, self._tunneledPort
)
self._protocol.transport.startTLS(sslOptions, self._protocolFactory)
self._tunnelReadyDeferred.callback(self._protocol)
else:
extra: Any
if respm:
extra = {
"status": int(respm.group("status")),
"reason": respm.group("reason").strip(),
}
else:
extra = rcvd_bytes[: self._truncatedLength]
extra = data[: self._truncatedLength]
self._tunnelReadyDeferred.errback(
TunnelError(
"Could not open CONNECT tunnel with proxy "
@ -187,11 +194,11 @@ class TunnelingTCP4ClientEndpoint(TCP4ClientEndpoint):
)
)
def connectFailed(self, reason):
def connectFailed(self, reason: Failure) -> None:
"""Propagates the errback to the appropriate deferred."""
self._tunnelReadyDeferred.errback(reason)
def connect(self, protocolFactory):
def connect(self, protocolFactory: Factory) -> Deferred:
self._protocolFactory = protocolFactory
connectDeferred = super().connect(protocolFactory)
connectDeferred.addCallback(self.requestTunnel)
@ -199,7 +206,9 @@ class TunnelingTCP4ClientEndpoint(TCP4ClientEndpoint):
return self._tunnelReadyDeferred
def tunnel_request_data(host, port, proxy_auth_header=None):
def tunnel_request_data(
host: str, port: int, proxy_auth_header: Optional[bytes] = None
) -> bytes:
r"""
Return binary content of a CONNECT request.
@ -230,18 +239,20 @@ class TunnelingAgent(Agent):
def __init__(
self,
reactor,
proxyConf,
contextFactory=None,
connectTimeout=None,
bindAddress=None,
pool=None,
reactor: ReactorBase,
proxyConf: Tuple[str, int, Optional[bytes]],
contextFactory: Optional[IPolicyForHTTPS] = None,
connectTimeout: Optional[float] = None,
bindAddress: Optional[bytes] = None,
pool: Optional[HTTPConnectionPool] = None,
):
# TODO make this arg required instead
assert contextFactory is not None
super().__init__(reactor, contextFactory, connectTimeout, bindAddress, pool)
self._proxyConf = proxyConf
self._contextFactory = contextFactory
self._proxyConf: Tuple[str, int, Optional[bytes]] = proxyConf
self._contextFactory: IPolicyForHTTPS = contextFactory
def _getEndpoint(self, uri):
def _getEndpoint(self, uri: URI) -> TunnelingTCP4ClientEndpoint:
return TunnelingTCP4ClientEndpoint(
reactor=self._reactor,
host=uri.host,
@ -253,8 +264,15 @@ class TunnelingAgent(Agent):
)
def _requestWithEndpoint(
self, key, endpoint, method, parsedURI, headers, bodyProducer, requestPath
):
self,
key: Any,
endpoint: TCP4ClientEndpoint,
method: bytes,
parsedURI: bytes,
headers: Optional[TxHeaders],
bodyProducer: Optional[IBodyProducer],
requestPath: bytes,
) -> Deferred:
# proxy host and port are required for HTTP pool `key`
# otherwise, same remote host connection request could reuse
# a cached tunneled connection to a different proxy
@ -272,7 +290,12 @@ class TunnelingAgent(Agent):
class ScrapyProxyAgent(Agent):
def __init__(
self, reactor, proxyURI, connectTimeout=None, bindAddress=None, pool=None
self,
reactor: ReactorBase,
proxyURI: bytes,
connectTimeout: Optional[float] = None,
bindAddress: Optional[bytes] = None,
pool: Optional[HTTPConnectionPool] = None,
):
super().__init__(
reactor=reactor,
@ -280,9 +303,15 @@ class ScrapyProxyAgent(Agent):
bindAddress=bindAddress,
pool=pool,
)
self._proxyURI = URI.fromBytes(proxyURI)
self._proxyURI: URI = URI.fromBytes(proxyURI)
def request(self, method, uri, headers=None, bodyProducer=None):
def request(
self,
method: bytes,
uri: bytes,
headers: Optional[TxHeaders] = None,
bodyProducer: Optional[IBodyProducer] = None,
) -> Deferred:
"""
Issue a new request via the configured proxy.
"""
@ -306,26 +335,29 @@ class ScrapyAgent:
def __init__(
self,
contextFactory=None,
connectTimeout=10,
bindAddress=None,
pool=None,
maxsize=0,
warnsize=0,
fail_on_dataloss=True,
crawler=None,
contextFactory: Optional[IPolicyForHTTPS] = None,
connectTimeout: float = 10,
bindAddress: Optional[bytes] = None,
pool: Optional[HTTPConnectionPool] = None,
maxsize: int = 0,
warnsize: int = 0,
fail_on_dataloss: bool = True,
crawler: Optional[Crawler] = None,
):
self._contextFactory = contextFactory
self._connectTimeout = connectTimeout
self._bindAddress = bindAddress
self._pool = pool
self._maxsize = maxsize
self._warnsize = warnsize
self._fail_on_dataloss = fail_on_dataloss
self._txresponse = None
self._crawler = crawler
# TODO make these args required instead
assert contextFactory is not None
assert crawler is not None
self._contextFactory: IPolicyForHTTPS = contextFactory
self._connectTimeout: float = connectTimeout
self._bindAddress: Optional[bytes] = bindAddress
self._pool: Optional[HTTPConnectionPool] = pool
self._maxsize: int = maxsize
self._warnsize: int = warnsize
self._fail_on_dataloss: bool = fail_on_dataloss
self._txresponse: Optional[TxResponse] = None
self._crawler: Crawler = crawler
def _get_agent(self, request, timeout):
def _get_agent(self, request: Request, timeout: float) -> Agent:
from twisted.internet import reactor
bindaddress = request.meta.get("bindaddress") or self._bindAddress
@ -333,10 +365,10 @@ class ScrapyAgent:
if proxy:
proxyScheme, proxyNetloc, proxyHost, proxyPort, proxyParams = _parse(proxy)
scheme = _parse(request.url)[0]
proxyHost = to_unicode(proxyHost)
proxyHost_str = to_unicode(proxyHost)
if scheme == b"https":
proxyAuth = request.headers.get(b"Proxy-Authorization", None)
proxyConf = (proxyHost, proxyPort, proxyAuth)
proxyConf = (proxyHost_str, proxyPort, proxyAuth)
return self._TunnelingAgent(
reactor=reactor,
proxyConf=proxyConf,
@ -346,7 +378,9 @@ class ScrapyAgent:
pool=self._pool,
)
proxyScheme = proxyScheme or b"http"
proxyURI = urlunparse((proxyScheme, proxyNetloc, proxyParams, "", "", ""))
proxyURI = urlunparse(
(proxyScheme, proxyNetloc, proxyParams, b"", b"", b"")
)
return self._ProxyAgent(
reactor=reactor,
proxyURI=to_bytes(proxyURI, encoding="ascii"),
@ -363,7 +397,7 @@ class ScrapyAgent:
pool=self._pool,
)
def download_request(self, request):
def download_request(self, request: Request) -> Deferred:
from twisted.internet import reactor
timeout = request.meta.get("download_timeout") or self._connectTimeout
@ -380,7 +414,7 @@ class ScrapyAgent:
else:
bodyproducer = None
start_time = time()
d = agent.request(
d: Deferred = agent.request(
method, to_bytes(url, encoding="ascii"), headers, bodyproducer
)
# set download latency
@ -393,7 +427,9 @@ class ScrapyAgent:
d.addBoth(self._cb_timeout, request, url, timeout)
return d
def _cb_timeout(self, result, request, url, timeout):
def _cb_timeout(
self, result: Any, request: Request, url: str, timeout: float
) -> Any:
if self._timeout_cl.active():
self._timeout_cl.cancel()
return result
@ -404,19 +440,21 @@ class ScrapyAgent:
raise TimeoutError(f"Getting {url} took longer than {timeout} seconds.")
def _cb_latency(self, result, request, start_time):
def _cb_latency(self, result: Any, request: Request, start_time: float) -> Any:
request.meta["download_latency"] = time() - start_time
return result
@staticmethod
def _headers_from_twisted_response(response):
def _headers_from_twisted_response(response: TxResponse) -> Headers:
headers = Headers()
if response.length != UNKNOWN_LENGTH:
headers[b"Content-Length"] = str(response.length).encode()
headers.update(response.headers.getAllRawHeaders())
return headers
def _cb_bodyready(self, txresponse, request):
def _cb_bodyready(
self, txresponse: TxResponse, request: Request
) -> Union[Dict[str, Any], Deferred]:
headers_received_result = self._crawler.signals.send_catch_log(
signal=signals.headers_received,
headers=self._headers_from_twisted_response(txresponse),
@ -472,7 +510,7 @@ class ScrapyAgent:
logger.warning(warning_msg, warning_args)
txresponse._transport.loseConnection()
raise defer.CancelledError(warning_msg % warning_args)
raise CancelledError(warning_msg % warning_args)
if warnsize and expected_size > warnsize:
logger.warning(
@ -481,11 +519,11 @@ class ScrapyAgent:
{"size": expected_size, "warnsize": warnsize, "request": request},
)
def _cancel(_):
def _cancel(_: Any) -> None:
# Abort connection immediately.
txresponse._transport._producer.abortConnection()
d = defer.Deferred(_cancel)
d: Deferred = Deferred(_cancel)
txresponse.deliverBody(
_ResponseReader(
finished=d,
@ -503,7 +541,9 @@ class ScrapyAgent:
return d
def _cb_bodydone(self, result, request, url):
def _cb_bodydone(
self, result: Dict[str, Any], request: Request, url: str
) -> Union[Response, Failure]:
headers = self._headers_from_twisted_response(result["txresponse"])
respcls = responsetypes.from_args(headers=headers, url=url, body=result["body"])
try:
@ -523,53 +563,57 @@ class ScrapyAgent:
)
if result.get("failure"):
result["failure"].value.response = response
return result["failure"]
return cast(Failure, result["failure"])
return response
@implementer(IBodyProducer)
class _RequestBodyProducer:
def __init__(self, body):
def __init__(self, body: bytes):
self.body = body
self.length = len(body)
def startProducing(self, consumer):
def startProducing(self, consumer: IConsumer) -> Deferred:
consumer.write(self.body)
return defer.succeed(None)
return succeed(None)
def pauseProducing(self):
def pauseProducing(self) -> None:
pass
def stopProducing(self):
def stopProducing(self) -> None:
pass
class _ResponseReader(protocol.Protocol):
class _ResponseReader(Protocol):
def __init__(
self,
finished,
txresponse,
request,
maxsize,
warnsize,
fail_on_dataloss,
crawler,
finished: Deferred,
txresponse: TxResponse,
request: Request,
maxsize: int,
warnsize: int,
fail_on_dataloss: bool,
crawler: Crawler,
):
self._finished = finished
self._txresponse = txresponse
self._request = request
self._bodybuf = BytesIO()
self._maxsize = maxsize
self._warnsize = warnsize
self._fail_on_dataloss = fail_on_dataloss
self._fail_on_dataloss_warned = False
self._reached_warnsize = False
self._bytes_received = 0
self._certificate = None
self._ip_address = None
self._crawler = crawler
self._finished: Deferred = finished
self._txresponse: TxResponse = txresponse
self._request: Request = request
self._bodybuf: BytesIO = BytesIO()
self._maxsize: int = maxsize
self._warnsize: int = warnsize
self._fail_on_dataloss: bool = fail_on_dataloss
self._fail_on_dataloss_warned: bool = False
self._reached_warnsize: bool = False
self._bytes_received: int = 0
self._certificate: Optional[ssl.Certificate] = None
self._ip_address: Union[ipaddress.IPv4Address, ipaddress.IPv6Address, None] = (
None
)
self._crawler: Crawler = crawler
def _finish_response(self, flags=None, failure=None):
def _finish_response(
self, flags: Optional[List[str]] = None, failure: Optional[Failure] = None
) -> None:
self._finished.callback(
{
"txresponse": self._txresponse,
@ -581,7 +625,8 @@ class _ResponseReader(protocol.Protocol):
}
)
def connectionMade(self):
def connectionMade(self) -> None:
assert self.transport
if self._certificate is None:
with suppress(AttributeError):
self._certificate = ssl.Certificate(
@ -593,11 +638,12 @@ class _ResponseReader(protocol.Protocol):
self.transport._producer.getPeer().host
)
def dataReceived(self, bodyBytes):
def dataReceived(self, bodyBytes: bytes) -> None:
# This maybe called several times after cancel was called with buffered data.
if self._finished.called:
return
assert self.transport
self._bodybuf.write(bodyBytes)
self._bytes_received += len(bodyBytes)
@ -644,7 +690,7 @@ class _ResponseReader(protocol.Protocol):
{"warnsize": self._warnsize, "request": self._request},
)
def connectionLost(self, reason):
def connectionLost(self, reason: Failure = connectionDone) -> None:
if self._finished.called:
return

View File

@ -8,6 +8,7 @@ from twisted.internet.base import DelayedCall
from twisted.internet.defer import Deferred
from twisted.internet.error import TimeoutError
from twisted.web.client import URI
from twisted.web.iweb import IPolicyForHTTPS
from scrapy.core.downloader.contextfactory import load_context_factory_from_settings
from scrapy.core.downloader.webclient import _parse
@ -24,7 +25,7 @@ if TYPE_CHECKING:
class H2DownloadHandler:
def __init__(self, settings: Settings, crawler: Optional[Crawler] = None):
def __init__(self, settings: Settings, crawler: Crawler):
self._crawler = crawler
from twisted.internet import reactor
@ -54,7 +55,7 @@ class ScrapyH2Agent:
def __init__(
self,
context_factory,
context_factory: IPolicyForHTTPS,
pool: H2ConnectionPool,
connect_timeout: int = 10,
bind_address: Optional[bytes] = None,

View File

@ -1,9 +1,14 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any, Optional, Type
from twisted.internet.defer import Deferred
from scrapy import Request, Spider
from scrapy.core.downloader.handlers.http import HTTPDownloadHandler
from scrapy.crawler import Crawler
from scrapy.exceptions import NotConfigured
from scrapy.settings import BaseSettings
from scrapy.utils.boto import is_botocore_available
from scrapy.utils.httpobj import urlparse_cached
from scrapy.utils.misc import build_from_crawler
@ -16,14 +21,14 @@ if TYPE_CHECKING:
class S3DownloadHandler:
def __init__(
self,
settings,
settings: BaseSettings,
*,
crawler=None,
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
httpdownloadhandler=HTTPDownloadHandler,
**kw,
crawler: Crawler,
aws_access_key_id: Optional[str] = None,
aws_secret_access_key: Optional[str] = None,
aws_session_token: Optional[str] = None,
httpdownloadhandler: Type[HTTPDownloadHandler] = HTTPDownloadHandler,
**kw: Any,
):
if not is_botocore_available():
raise NotConfigured("missing botocore library")
@ -51,6 +56,8 @@ class S3DownloadHandler:
if kw:
raise TypeError(f"Unexpected keyword arguments: {kw}")
if not self.anon:
assert aws_access_key_id is not None
assert aws_secret_access_key is not None
SignerCls = botocore.auth.AUTH_TYPE_MAPS["s3"]
self._signer = SignerCls(
botocore.credentials.Credentials(
@ -65,10 +72,10 @@ class S3DownloadHandler:
self._download_http = _http_handler.download_request
@classmethod
def from_crawler(cls, crawler, **kwargs) -> Self:
def from_crawler(cls, crawler: Crawler, **kwargs: Any) -> Self:
return cls(crawler.settings, crawler=crawler, **kwargs)
def download_request(self, request, spider):
def download_request(self, request: Request, spider: Spider) -> Deferred:
p = urlparse_cached(request)
scheme = "https" if request.meta.get("is_secure") else "http"
bucket = p.hostname
@ -85,6 +92,7 @@ class S3DownloadHandler:
headers=request.headers.to_unicode_dict(),
data=request.body,
)
assert self._signer
self._signer.add_auth(awsrequest)
request = request.replace(url=url, headers=awsrequest.headers.items())
return self._download_http(request, spider)