mirror of https://github.com/scrapy/scrapy.git
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:
parent
031bfc9c3b
commit
92bec38591
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in New Issue