diff --git a/scrapy/core/http2/protocol.py b/scrapy/core/http2/protocol.py index a3dfdb76e..2f177656d 100644 --- a/scrapy/core/http2/protocol.py +++ b/scrapy/core/http2/protocol.py @@ -11,32 +11,34 @@ from h2.events import ( StreamEnded, StreamReset, WindowUpdated ) from h2.exceptions import ProtocolError +from twisted.internet.defer import Deferred from twisted.internet.protocol import connectionDone, Protocol from twisted.internet.ssl import Certificate +from twisted.python.failure import Failure from scrapy.core.http2.stream import Stream, StreamCloseReason from scrapy.core.http2.types import H2ConnectionMetadataDict from scrapy.http import Request + logger = logging.getLogger(__name__) class H2ClientProtocol(Protocol): - def __init__(self): + def __init__(self) -> None: config = H2Configuration(client_side=True, header_encoding='utf-8') self.conn = H2Connection(config=config) # ID of the next request stream - # Following the convention made by hyper-h2 each client ID - # will be odd. - self.stream_id_count = itertools.count(start=1, step=2) + # Following the convention made by hyper-h2 all IDs will be odd + self._stream_id_generator = itertools.count(start=1, step=2) # Streams are stored in a dictionary keyed off their stream IDs self.streams: Dict[int, Stream] = {} # If requests are received before connection is made we keep # all requests in a pool and send them as the connection is made - self._pending_request_stream_pool = deque() + self._pending_request_stream_pool: deque = deque() # Counter to keep track of opened stream. This counter # is used to make sure that not more than MAX_CONCURRENT_STREAMS @@ -48,15 +50,15 @@ class H2ClientProtocol(Protocol): # We pass this instance to the streams ResponseFailed() failure self._protocol_error: Optional[ProtocolError] = None - self._metadata: H2ConnectionMetadataDict = { + self.metadata: H2ConnectionMetadataDict = { 'certificate': None, 'ip_address': None, 'hostname': None, - 'port': None + 'port': None, } @property - def is_connected(self): + def is_connected(self) -> bool: """Boolean to keep track of the connection status. This is used while initiating pending streams to make sure that we initiate stream only during active HTTP/2 Connection @@ -75,7 +77,7 @@ class H2ClientProtocol(Protocol): self.conn.remote_settings.max_concurrent_streams ) - def _send_pending_requests(self): + def _send_pending_requests(self) -> None: """Initiate all pending requests from the deque following FIFO We make sure that at any time {allowed_max_concurrent_streams} streams are active. @@ -89,37 +91,33 @@ class H2ClientProtocol(Protocol): stream = self._pending_request_stream_pool.popleft() stream.initiate_request() - def _stream_close_cb(self, stream_id: int): - """Called when stream is closed completely + def pop_stream(self, stream_id: int) -> Stream: + """Perform cleanup when a stream is closed """ - self.streams.pop(stream_id) + stream = self.streams.pop(stream_id) self._active_streams -= 1 self._send_pending_requests() + return stream - def _new_stream(self, request: Request): + def _new_stream(self, request: Request) -> Stream: """Instantiates a new Stream object """ - stream_id = next(self.stream_id_count) - stream = Stream( - stream_id=stream_id, + stream_id=next(self._stream_id_generator), request=request, - connection=self.conn, - conn_metadata=self._metadata, - cb_close=self._stream_close_cb + protocol=self, ) - self.streams[stream.stream_id] = stream return stream - def _write_to_transport(self): + def _write_to_transport(self) -> None: """ Write data to the underlying transport connection from the HTTP2 connection instance if any """ data = self.conn.data_to_send() self.transport.write(data) - def request(self, request: Request): + def request(self, request: Request) -> Deferred: if not isinstance(request, Request): raise TypeError(f'Expected scrapy.http.Request, received {request.__class__.__qualname__}') @@ -130,20 +128,20 @@ class H2ClientProtocol(Protocol): self._pending_request_stream_pool.append(stream) return d - def connectionMade(self): + def connectionMade(self) -> None: """Called by Twisted when the connection is established. We can start sending some data now: we should open with the connection preamble. """ destination = self.transport.getPeer() logger.debug('Connection made to {}'.format(destination)) - self._metadata['ip_address'] = ipaddress.ip_address(destination.host) - self._metadata['port'] = destination.port - self._metadata['hostname'] = self.transport.transport.addr[0] + self.metadata['ip_address'] = ipaddress.ip_address(destination.host) + self.metadata['port'] = destination.port + self.metadata['hostname'] = self.transport.transport.addr[0] self.conn.initiate_connection() self._write_to_transport() - def dataReceived(self, data): + def dataReceived(self, data: bytes) -> None: try: events = self.conn.receive_data(data) self._handle_events(events) @@ -158,32 +156,30 @@ class H2ClientProtocol(Protocol): finally: self._write_to_transport() - def connectionLost(self, reason=connectionDone): + def connectionLost(self, reason: Failure = connectionDone) -> None: """Called by Twisted when the transport connection is lost. No need to write anything to transport here. """ - # Pop all streams which were pending and were not yet started - # NOTE: Stream.close() pops the element from the streams dictionary - # which raises `RuntimeError: dictionary changed size during iteration` - # Hence, we copy the streams into a list. - for stream in list(self.streams.values()): + for stream in self.streams.values(): if stream.request_sent: - stream.close(StreamCloseReason.CONNECTION_LOST, self._protocol_error) + stream.close(StreamCloseReason.CONNECTION_LOST, self._protocol_error, from_protocol=True) else: - stream.close(StreamCloseReason.INACTIVE) + stream.close(StreamCloseReason.INACTIVE, from_protocol=True) + self._active_streams -= len(self.streams) + self.streams.clear() + self._send_pending_requests() self.conn.close_connection() if not reason.check(connectionDone): logger.warning("Connection lost with reason " + str(reason)) - def _handle_events(self, events): + def _handle_events(self, events: list) -> None: """Private method which acts as a bridge between the events received from the HTTP/2 data and IH2EventsHandler Arguments: - events {list} -- A list of events that the remote peer - triggered by sending data + events -- A list of events that the remote peer triggered by sending data """ for event in events: if isinstance(event, DataReceived): @@ -202,27 +198,29 @@ class H2ClientProtocol(Protocol): logger.debug('Received unhandled event {}'.format(event)) # Event handler functions starts here - def data_received(self, event: DataReceived): + def data_received(self, event: DataReceived) -> None: self.streams[event.stream_id].receive_data(event.data, event.flow_controlled_length) - def response_received(self, event: ResponseReceived): + def response_received(self, event: ResponseReceived) -> None: self.streams[event.stream_id].receive_headers(event.headers) - def settings_acknowledged(self, event: SettingsAcknowledged): + def settings_acknowledged(self, event: SettingsAcknowledged) -> None: # Send off all the pending requests as now we have # established a proper HTTP/2 connection self._send_pending_requests() # Update certificate when our HTTP/2 connection is established - self._metadata['certificate'] = Certificate(self.transport.getPeerCertificate()) + self.metadata['certificate'] = Certificate(self.transport.getPeerCertificate()) - def stream_ended(self, event: StreamEnded): - self.streams[event.stream_id].close(StreamCloseReason.ENDED) + def stream_ended(self, event: StreamEnded) -> None: + stream = self.pop_stream(event.stream_id) + stream.close(StreamCloseReason.ENDED, from_protocol=True) - def stream_reset(self, event: StreamReset): - self.streams[event.stream_id].close(StreamCloseReason.RESET) + def stream_reset(self, event: StreamReset) -> None: + stream = self.pop_stream(event.stream_id) + stream.close(StreamCloseReason.RESET, from_protocol=True) - def window_updated(self, event: WindowUpdated): + def window_updated(self, event: WindowUpdated) -> None: if event.stream_id != 0: self.streams[event.stream_id].receive_window_update() else: diff --git a/scrapy/core/http2/stream.py b/scrapy/core/http2/stream.py index 1c856ff68..77cfbcfbf 100644 --- a/scrapy/core/http2/stream.py +++ b/scrapy/core/http2/stream.py @@ -1,22 +1,27 @@ import logging from enum import Enum from io import BytesIO -from typing import Callable, List +from typing import List, Optional, Tuple, TYPE_CHECKING from urllib.parse import urlparse -from h2.connection import H2Connection from h2.errors import ErrorCodes from h2.exceptions import StreamClosedError +from hpack import HeaderTuple from twisted.internet.defer import Deferred, CancelledError from twisted.internet.error import ConnectionClosed from twisted.python.failure import Failure from twisted.web.client import ResponseFailed -from scrapy.core.http2.types import H2ConnectionMetadataDict, H2ResponseDict +from scrapy.core.http2.types import H2ResponseDict from scrapy.http import Request from scrapy.http.headers import Headers from scrapy.responsetypes import responsetypes + +if TYPE_CHECKING: + from scrapy.core.http2.protocol import H2ClientProtocol + + logger = logging.getLogger(__name__) @@ -31,12 +36,12 @@ class InactiveStreamClosed(ConnectionClosed): class InvalidHostname(Exception): - def __init__(self, request: Request, expected_hostname, expected_netloc): + def __init__(self, request: Request, expected_hostname: Optional[str], expected_netloc: Optional[str]) -> None: self.request = request self.expected_hostname = expected_hostname self.expected_netloc = expected_netloc - def __str__(self): + def __str__(self) -> str: return f'InvalidHostname: Expected {self.expected_hostname} or {self.expected_netloc} in {self.request}' @@ -80,28 +85,20 @@ class Stream: self, stream_id: int, request: Request, - connection: H2Connection, - conn_metadata: H2ConnectionMetadataDict, - cb_close: Callable[[int], None], + protocol: "H2ClientProtocol", download_maxsize: int = 0, download_warnsize: int = 0, fail_on_data_loss: bool = True - ): + ) -> None: """ Arguments: - stream_id -- For one HTTP/2 connection each stream is - uniquely identified by a single integer - request -- HTTP request - connection -- HTTP/2 connection this stream belongs to. - conn_metadata -- Reference to dictionary having metadata of HTTP/2 connection - cb_close -- Method called when this stream is closed - to notify the TCP connection instance. + stream_id -- Unique identifier for the stream within a single HTTP/2 connection + request -- The HTTP request associated to the stream + protocol -- Parent H2ClientProtocol instance """ - self.stream_id = stream_id - self._request = request - self._conn = connection - self._conn_metadata = conn_metadata - self._cb_close = cb_close + self.stream_id: int = stream_id + self._request: Request = request + self._protocol: "H2ClientProtocol" = protocol self._download_maxsize = self._request.meta.get('download_maxsize', download_maxsize) self._download_warnsize = self._request.meta.get('download_warnsize', download_warnsize) @@ -132,7 +129,7 @@ class Stream: self._response: H2ResponseDict = { 'body': BytesIO(), 'flow_controlled_size': 0, - 'headers': Headers({}) + 'headers': Headers({}), } def _cancel(_): @@ -145,7 +142,7 @@ class Stream: self._deferred_response = Deferred(_cancel) - def __str__(self): + def __str__(self) -> str: return f'Stream(id={self.stream_id!r})' __repr__ = __str__ @@ -169,12 +166,9 @@ class Stream: and not self._reached_warnsize ) - def get_response(self): + def get_response(self) -> Deferred: """Simply return a Deferred which fires when response from the asynchronous request is available - - Returns: - Deferred -- Calls the callback passing the response """ return self._deferred_response @@ -182,12 +176,12 @@ class Stream: # Make sure that we are sending the request to the correct URL url = urlparse(self._request.url) return ( - url.netloc == self._conn_metadata['hostname'] - or url.netloc == f'{self._conn_metadata["hostname"]}:{self._conn_metadata["port"]}' - or url.netloc == f'{self._conn_metadata["ip_address"]}:{self._conn_metadata["port"]}' + url.netloc == self._protocol.metadata['hostname'] + or url.netloc == f'{self._protocol.metadata["hostname"]}:{self._protocol.metadata["port"]}' + or url.netloc == f'{self._protocol.metadata["ip_address"]}:{self._protocol.metadata["port"]}' ) - def _get_request_headers(self): + def _get_request_headers(self) -> List[Tuple[str, str]]: url = urlparse(self._request.url) path = url.path @@ -207,10 +201,10 @@ class Stream: return headers - def initiate_request(self): + def initiate_request(self) -> None: if self.check_request_url(): headers = self._get_request_headers() - self._conn.send_headers(self.stream_id, headers, end_stream=False) + self._protocol.conn.send_headers(self.stream_id, headers, end_stream=False) self.request_sent = True self.send_data() else: @@ -218,7 +212,7 @@ class Stream: # Note that we have not sent any headers self.close(StreamCloseReason.INVALID_HOSTNAME) - def send_data(self): + def send_data(self) -> None: """Called immediately after the headers are sent. Here we send all the data as part of the request. @@ -233,10 +227,10 @@ class Stream: raise StreamClosedError(self.stream_id) # Firstly, check what the flow control window is for current stream. - window_size = self._conn.local_flow_control_window(stream_id=self.stream_id) + window_size = self._protocol.conn.local_flow_control_window(stream_id=self.stream_id) # Next, check what the maximum frame size is. - max_frame_size = self._conn.max_outbound_frame_size + max_frame_size = self._protocol.conn.max_outbound_frame_size # We will send no more than the window size or the remaining file size # of data in this call, whichever is smaller. @@ -249,7 +243,7 @@ class Stream: data_chunk_start_id = self.content_length - self.remaining_content_length data_chunk = self._request.body[data_chunk_start_id:data_chunk_start_id + chunk_size] - self._conn.send_data(self.stream_id, data_chunk, end_stream=False) + self._protocol.conn.send_data(self.stream_id, data_chunk, end_stream=False) bytes_to_send_size = bytes_to_send_size - chunk_size self.remaining_content_length = self.remaining_content_length - chunk_size @@ -258,12 +252,12 @@ class Stream: # End the stream if no more data needs to be send if self.remaining_content_length == 0: - self._conn.end_stream(self.stream_id) + self._protocol.conn.end_stream(self.stream_id) # Q. What about the rest of the data? # Ans: Remaining Data frames will be sent when we get a WindowUpdate frame - def receive_window_update(self): + def receive_window_update(self) -> None: """Flow control window size was changed. Send data that earlier could not be sent as we were blocked behind the flow control. @@ -271,7 +265,7 @@ class Stream: if self.remaining_content_length and not self.stream_closed_server and self.request_sent: self.send_data() - def receive_data(self, data: bytes, flow_controlled_length: int): + def receive_data(self, data: bytes, flow_controlled_length: int) -> None: self._response['body'].write(data) self._response['flow_controlled_size'] += flow_controlled_length @@ -289,12 +283,12 @@ class Stream: logger.warning(warning_msg) # Acknowledge the data received - self._conn.acknowledge_received_data( + self._protocol.conn.acknowledge_received_data( self._response['flow_controlled_size'], self.stream_id ) - def receive_headers(self, headers): + def receive_headers(self, headers: List[HeaderTuple]) -> None: for name, value in headers: self._response['headers'][name] = value @@ -312,7 +306,7 @@ class Stream: ) logger.warning(warning_msg) - def reset_stream(self, reason=StreamCloseReason.RESET): + def reset_stream(self, reason: StreamCloseReason = StreamCloseReason.RESET) -> None: """Close this stream by sending a RST_FRAME to the remote peer""" if self.stream_closed_local: raise StreamClosedError(self.stream_id) @@ -321,7 +315,7 @@ class Stream: self._response['body'].truncate(0) self.stream_closed_local = True - self._conn.reset_stream(self.stream_id, ErrorCodes.REFUSED_STREAM) + self._protocol.conn.reset_stream(self.stream_id, ErrorCodes.REFUSED_STREAM) self.close(reason) def _is_data_lost(self) -> bool: @@ -332,11 +326,8 @@ class Stream: return expected_size != received_body_size - def close(self, reason: StreamCloseReason, error: Exception = None): + def close(self, reason: StreamCloseReason, error: Optional[Exception] = None, from_protocol: bool = False) -> None: """Based on the reason sent we will handle each case. - - Arguments: - reason -- One if StreamCloseReason """ if self.stream_closed_server: raise StreamClosedError(self.stream_id) @@ -344,7 +335,9 @@ class Stream: if not isinstance(reason, StreamCloseReason): raise TypeError(f'Expected StreamCloseReason, received {reason.__class__.__qualname__}') - self._cb_close(self.stream_id) + if not from_protocol: + self._protocol.pop_stream(self.stream_id) + self.stream_closed_server = True flags = None @@ -392,11 +385,11 @@ class Stream: elif reason is StreamCloseReason.INVALID_HOSTNAME: self._deferred_response.errback(InvalidHostname( self._request, - self._conn_metadata['hostname'], - f'{self._conn_metadata["ip_address"]}:{self._conn_metadata["port"]}' + self._protocol.metadata['hostname'], + f'{self._protocol.metadata["ip_address"]}:{self._protocol.metadata["port"]}' )) - def _fire_response_deferred(self, flags: List[str] = None): + def _fire_response_deferred(self, flags: Optional[List[str]] = None) -> None: """Builds response from the self._response dict and fires the response deferred callback with the generated response instance""" @@ -405,7 +398,7 @@ class Stream: response_cls = responsetypes.from_args( headers=self._response['headers'], url=self._request.url, - body=body + body=body, ) response = response_cls( @@ -415,8 +408,8 @@ class Stream: body=body, request=self._request, flags=flags, - certificate=self._conn_metadata['certificate'], - ip_address=self._conn_metadata['ip_address'] + certificate=self._protocol.metadata['certificate'], + ip_address=self._protocol.metadata['ip_address'], ) self._deferred_response.callback(response)