feat: MethodNotAllowed405, Content-Length header

- add tests to check for Content-Length header
- raise MethodNotAllowed405 when remote send 'HTTP/2.0 405 Method Not
Allowed'
This commit is contained in:
Aditya 2020-07-29 13:43:59 +05:30
parent 031bfc9c3b
commit 92bec38591
5 changed files with 85 additions and 24 deletions

View File

@ -3,6 +3,7 @@ from time import time
from typing import Optional, Tuple
from urllib.parse import urldefrag
from twisted.internet.defer import Deferred
from twisted.internet.base import ReactorBase
from twisted.internet.error import TimeoutError
from twisted.web.client import URI
@ -23,9 +24,6 @@ class H2DownloadHandler:
from twisted.internet import reactor
self._pool = H2ConnectionPool(reactor, settings)
self._context_factory = 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')
@classmethod
def from_crawler(cls, crawler):
@ -35,8 +33,6 @@ class H2DownloadHandler:
agent = ScrapyH2Agent(
context_factory=self._context_factory,
pool=self._pool,
maxsize=getattr(spider, 'download_maxsize', self._default_maxsize),
warnsize=getattr(spider, 'download_warnsize', self._default_warnsize),
crawler=self._crawler
)
return agent.download_request(request, spider)
@ -72,15 +68,12 @@ class ScrapyH2Agent:
self, context_factory,
connect_timeout=10,
bind_address: Optional[bytes] = None, pool: H2ConnectionPool = None,
maxsize: int = 0, warnsize: int = 0,
crawler=None
) -> None:
self._context_factory = context_factory
self._connect_timeout = connect_timeout
self._bind_address = bind_address
self._pool = pool
self._maxsize = maxsize
self._warnsize = warnsize
self._crawler = crawler
def _get_agent(self, request: Request, timeout: Optional[float]) -> H2Agent:
@ -121,7 +114,7 @@ class ScrapyH2Agent:
pool=self._pool
)
def download_request(self, request: Request, spider: Spider):
def download_request(self, request: Request, spider: Spider) -> Deferred:
from twisted.internet import reactor
timeout = request.meta.get('download_timeout') or self._connect_timeout
agent = self._get_agent(request, timeout)
@ -135,12 +128,12 @@ class ScrapyH2Agent:
return d
@staticmethod
def _cb_latency(response: Response, request: Request, start_time: float):
def _cb_latency(response: Response, request: Request, start_time: float) -> Response:
request.meta['download_latency'] = time() - start_time
return response
@staticmethod
def _cb_timeout(response: Response, request: Request, timeout: float, timeout_cl):
def _cb_timeout(response: Response, request: Request, timeout: float, timeout_cl) -> Response:
if timeout_cl.active():
timeout_cl.cancel()
return response

View File

@ -93,7 +93,7 @@ class H2ConnectionPool:
Deferred that fires when all connections have been closed
"""
for conn in self._connections.values():
conn.transport.loseConnection()
conn.transport.abortConnection()
@implementer(IPolicyForHTTPS)

View File

@ -13,7 +13,7 @@ from h2.events import (
SettingsAcknowledged, StreamEnded, StreamReset, UnknownFrameReceived,
WindowUpdated
)
from h2.exceptions import H2Error, ProtocolError
from h2.exceptions import H2Error
from twisted.internet.defer import Deferred
from twisted.internet.error import TimeoutError
from twisted.internet.interfaces import IHandshakeListener, IProtocolNegotiationFactory
@ -39,7 +39,7 @@ class InvalidNegotiatedProtocol(H2Error):
self.negotiated_protocol = negotiated_protocol
def __str__(self) -> str:
return f'InvalidHostname: Expected h2 as negotiated protocol, received {self.negotiated_protocol}'
return f'InvalidHostname: Expected h2 as negotiated protocol, received {self.negotiated_protocol!r}'
class RemoteTerminatedConnection(H2Error):
@ -48,7 +48,15 @@ class RemoteTerminatedConnection(H2Error):
self.terminate_event = event
def __str__(self) -> str:
return f'RemoteTerminatedConnection: Received GOAWAY frame from {self.remote_ip_address}'
return f'RemoteTerminatedConnection: Received GOAWAY frame from {self.remote_ip_address!r}'
class MethodNotAllowed405(H2Error):
def __init__(self, remote_ip_address: Optional[Union[IPv4Address, IPv6Address]]):
self.remote_ip_address = remote_ip_address
def __str__(self) -> str:
return f"MethodNotAllowed405: Received 'HTTP/2.0 405 Method Not Allowed' from {self.remote_ip_address!r}"
@implementer(IHandshakeListener)
@ -217,14 +225,25 @@ class H2ClientProtocol(Protocol, TimeoutMixin):
# So, no need to send a GOAWAY frame to the remote
self._lose_connection_with_error([InvalidNegotiatedProtocol(negotiated_protocol)])
def _check_received_data(self, data: bytes) -> None:
"""Checks for edge cases where the connection to remote fails
without raising an appropriate H2Error
Arguments:
data -- Data received from the remote
"""
if data.startswith(b'HTTP/2.0 405 Method Not Allowed'):
raise MethodNotAllowed405(self.metadata['ip_address'])
def dataReceived(self, data: bytes) -> None:
# Reset the idle timeout as connection is still actively receiving data
self.resetTimeout()
try:
self._check_received_data(data)
events = self.conn.receive_data(data)
self._handle_events(events)
except ProtocolError as e:
except H2Error as e:
# Save this error as ultimately the connection will be dropped
# internally by hyper-h2. Saved error will be passed to all the streams
# closed with the connection.
@ -271,9 +290,10 @@ class H2ClientProtocol(Protocol, TimeoutMixin):
for stream in self.streams.values():
if stream.request_sent:
stream.close(StreamCloseReason.CONNECTION_LOST, self._conn_lost_errors, from_protocol=True)
close_reason = StreamCloseReason.CONNECTION_LOST
else:
stream.close(StreamCloseReason.INACTIVE, from_protocol=True)
close_reason = StreamCloseReason.INACTIVE
stream.close(close_reason, self._conn_lost_errors, from_protocol=True)
self._active_streams -= len(self.streams)
self.streams.clear()

View File

@ -189,12 +189,15 @@ class Stream:
headers = [
(':method', self._request.method),
(':authority', url.netloc),
(':scheme', 'https'),
(':scheme', self._protocol.metadata['uri'].scheme),
(':path', path),
]
for name, value in self._request.headers.items():
headers.append((name, value[0]))
headers.append((str(name, 'utf-8'), str(value[0], 'utf-8')))
if b'Content-Length' not in self._request.headers.keys():
headers.append(('Content-Length', str(len(self._request.body))))
return headers
@ -337,6 +340,10 @@ class Stream:
if not isinstance(reason, StreamCloseReason):
raise TypeError(f'Expected StreamCloseReason, received {reason.__class__.__qualname__}')
# Have default value of errors as an empty list as
# some cases can add a list of exceptions
errors = errors or []
if not from_protocol:
self._protocol.pop_stream(self.stream_id)
@ -387,7 +394,8 @@ class Stream:
self._deferred_response.errback(ResponseFailed(errors))
elif reason is StreamCloseReason.INACTIVE:
self._deferred_response.errback(InactiveStreamClosed(self._request))
errors.insert(0, InactiveStreamClosed(self._request))
self._deferred_response.errback(ResponseFailed(errors))
elif reason is StreamCloseReason.INVALID_HOSTNAME:
self._deferred_response.errback(InvalidHostname(

View File

@ -11,14 +11,14 @@ from h2.exceptions import InvalidBodyLengthError
from twisted.internet import reactor
from twisted.internet.defer import CancelledError, Deferred, DeferredList, inlineCallbacks
from twisted.internet.endpoints import SSL4ClientEndpoint, SSL4ServerEndpoint
from twisted.internet.error import TimeoutError
from twisted.internet.ssl import optionsForClientTLS, PrivateCertificate, Certificate
from twisted.python.failure import Failure
from twisted.trial.unittest import TestCase
from twisted.web.client import URI
from twisted.web.client import ResponseFailed, URI
from twisted.web.http import Request as TxRequest
from twisted.web.server import Site, NOT_DONE_YET
from twisted.web.static import File
from twisted.internet.error import TimeoutError
from scrapy.core.http2.protocol import H2ClientFactory, H2ClientProtocol
from scrapy.core.http2.stream import InactiveStreamClosed, InvalidHostname
@ -152,6 +152,19 @@ class QueryParams(LeafResource):
return bytes(json.dumps(query_params), 'utf-8')
class RequestHeaders(LeafResource):
"""Sends all the headers received as a response"""
def render_GET(self, request: TxRequest):
request.setHeader('Content-Type', 'application/json; charset=UTF-8')
request.setHeader('Content-Encoding', 'UTF-8')
headers = {}
for k, v in request.requestHeaders.getAllRawHeaders():
headers[str(k, 'utf-8')] = str(v[0], 'utf-8')
return bytes(json.dumps(headers), 'utf-8')
def get_client_certificate(key_file, certificate_file) -> PrivateCertificate:
with open(key_file, 'r') as key, open(certificate_file, 'r') as certificate:
pem = ''.join(key.readlines()) + ''.join(certificate.readlines())
@ -179,6 +192,7 @@ class Https2ClientProtocolTestCase(TestCase):
r.putChild(b'status', Status())
r.putChild(b'query-params', QueryParams())
r.putChild(b'timeout', TimeoutResponse())
r.putChild(b'request-headers', RequestHeaders())
return r
@inlineCallbacks
@ -488,7 +502,11 @@ class Https2ClientProtocolTestCase(TestCase):
d_list = []
def assert_inactive_stream(failure):
self.assertIsNotNone(failure.check(InactiveStreamClosed))
self.assertIsNotNone(failure.check(ResponseFailed))
self.assertTrue(any(
isinstance(e, InactiveStreamClosed)
for e in failure.value.reasons
))
# Send 100 request (we do not check the result)
for _ in range(100):
@ -616,3 +634,25 @@ class Https2ClientProtocolTestCase(TestCase):
d.addCallback(self.fail)
d.addErrback(assert_timeout_error)
return d
def test_request_headers_received(self):
request = Request(self.get_url('/request-headers'), headers={
'header-1': 'header value 1',
'header-2': 'header value 2'
})
d = self.make_request(request)
def assert_request_headers(response: Response):
self.assertEqual(response.status, 200)
self.assertEqual(response.request, request)
response_headers = json.loads(str(response.body, 'utf-8'))
self.assertIsInstance(response_headers, dict)
for k, v in request.headers.items():
k, v = str(k, 'utf-8'), str(v[0], 'utf-8')
self.assertIn(k, response_headers)
self.assertEqual(v, response_headers[k])
d.addErrback(self.fail)
d.addCallback(assert_request_headers)
return d