mirror of https://github.com/scrapy/scrapy.git
194 lines
7.2 KiB
Python
194 lines
7.2 KiB
Python
from __future__ import annotations
|
|
|
|
import warnings
|
|
from typing import TYPE_CHECKING, Any, cast
|
|
|
|
import OpenSSL.SSL
|
|
import pytest
|
|
from pytest_twisted import async_yield_fixture
|
|
from twisted.web import server, static
|
|
from twisted.web.client import Agent, BrowserLikePolicyForHTTPS, readBody
|
|
from twisted.web.client import Response as TxResponse
|
|
|
|
from scrapy.core.downloader import Downloader, Slot
|
|
from scrapy.core.downloader.contextfactory import (
|
|
ScrapyClientContextFactory,
|
|
load_context_factory_from_settings,
|
|
)
|
|
from scrapy.core.downloader.handlers.http11 import _RequestBodyProducer
|
|
from scrapy.exceptions import ScrapyDeprecationWarning
|
|
from scrapy.settings import Settings
|
|
from scrapy.utils.defer import maybe_deferred_to_future
|
|
from scrapy.utils.misc import build_from_crawler
|
|
from scrapy.utils.python import to_bytes
|
|
from scrapy.utils.spider import DefaultSpider
|
|
from scrapy.utils.test import get_crawler
|
|
from tests.mockserver.http_resources import PayloadResource
|
|
from tests.mockserver.utils import ssl_context_factory
|
|
from tests.utils.decorators import coroutine_test
|
|
|
|
if TYPE_CHECKING:
|
|
from twisted.internet.defer import Deferred
|
|
from twisted.internet.ssl import ContextFactory
|
|
from twisted.web.iweb import IBodyProducer
|
|
|
|
|
|
class TestSlot:
|
|
def test_repr(self):
|
|
slot = Slot(concurrency=8, delay=0.1, randomize_delay=True)
|
|
assert repr(slot) == "Slot(concurrency=8, delay=0.10, randomize_delay=True)"
|
|
|
|
|
|
@pytest.mark.requires_reactor # this test is related to the Twisted HTTP code
|
|
class TestContextFactoryBase:
|
|
context_factory: ContextFactory | None = None
|
|
|
|
@async_yield_fixture
|
|
async def server_url(self, tmp_path):
|
|
(tmp_path / "file").write_bytes(b"0123456789")
|
|
r = static.File(str(tmp_path))
|
|
r.putChild(b"payload", PayloadResource())
|
|
site = server.Site(r, timeout=None)
|
|
port = self._listen(site)
|
|
portno = port.getHost().port
|
|
|
|
yield f"https://127.0.0.1:{portno}/"
|
|
|
|
await port.stopListening()
|
|
|
|
def _listen(self, site):
|
|
from twisted.internet import reactor
|
|
|
|
return reactor.listenSSL(
|
|
0,
|
|
site,
|
|
contextFactory=self.context_factory or ssl_context_factory(),
|
|
interface="127.0.0.1",
|
|
)
|
|
|
|
@staticmethod
|
|
async def get_page(
|
|
url: str,
|
|
client_context_factory: BrowserLikePolicyForHTTPS,
|
|
body: str | None = None,
|
|
) -> bytes:
|
|
from twisted.internet import reactor
|
|
|
|
agent = Agent(reactor, contextFactory=client_context_factory)
|
|
body_producer = _RequestBodyProducer(body.encode()) if body else None
|
|
response: TxResponse = cast(
|
|
"TxResponse",
|
|
await maybe_deferred_to_future(
|
|
agent.request(
|
|
b"GET",
|
|
url.encode(),
|
|
bodyProducer=cast("IBodyProducer", body_producer),
|
|
)
|
|
),
|
|
)
|
|
with warnings.catch_warnings():
|
|
# https://github.com/twisted/twisted/issues/8227
|
|
warnings.filterwarnings(
|
|
"ignore",
|
|
category=DeprecationWarning,
|
|
message=r".*does not have an abortConnection method",
|
|
)
|
|
d: Deferred[bytes] = readBody(response) # type: ignore[arg-type]
|
|
return await maybe_deferred_to_future(d)
|
|
|
|
|
|
class TestContextFactory(TestContextFactoryBase):
|
|
@coroutine_test
|
|
async def testPayload(self, server_url: str) -> None:
|
|
s = "0123456789" * 10
|
|
crawler = get_crawler()
|
|
settings = Settings()
|
|
client_context_factory = load_context_factory_from_settings(settings, crawler)
|
|
body = await self.get_page(
|
|
server_url + "payload", client_context_factory, body=s
|
|
)
|
|
assert body == to_bytes(s)
|
|
|
|
def test_override_getContext(self):
|
|
class MyFactory(ScrapyClientContextFactory):
|
|
def getContext(
|
|
self, hostname: Any = None, port: Any = None
|
|
) -> OpenSSL.SSL.Context:
|
|
ctx: OpenSSL.SSL.Context = super().getContext(hostname, port)
|
|
return ctx
|
|
|
|
with warnings.catch_warnings(record=True) as w:
|
|
MyFactory()
|
|
assert len(w) == 1
|
|
assert (
|
|
"Overriding ScrapyClientContextFactory.getContext() is deprecated"
|
|
in str(w[0].message)
|
|
)
|
|
|
|
|
|
class TestContextFactoryTLSMethod(TestContextFactoryBase):
|
|
async def _assert_factory_works(
|
|
self, server_url: str, client_context_factory: ScrapyClientContextFactory
|
|
) -> None:
|
|
s = "0123456789" * 10
|
|
body = await self.get_page(
|
|
server_url + "payload", client_context_factory, body=s
|
|
)
|
|
assert body == to_bytes(s)
|
|
|
|
@coroutine_test
|
|
async def test_setting_default(self, server_url: str) -> None:
|
|
crawler = get_crawler()
|
|
settings = Settings()
|
|
client_context_factory = load_context_factory_from_settings(settings, crawler)
|
|
assert client_context_factory._ssl_method == OpenSSL.SSL.SSLv23_METHOD
|
|
await self._assert_factory_works(server_url, client_context_factory)
|
|
|
|
def test_setting_none(self):
|
|
crawler = get_crawler()
|
|
settings = Settings({"DOWNLOADER_CLIENT_TLS_METHOD": None})
|
|
with pytest.raises(KeyError):
|
|
load_context_factory_from_settings(settings, crawler)
|
|
|
|
def test_setting_bad(self):
|
|
crawler = get_crawler()
|
|
settings = Settings({"DOWNLOADER_CLIENT_TLS_METHOD": "bad"})
|
|
with pytest.raises(KeyError):
|
|
load_context_factory_from_settings(settings, crawler)
|
|
|
|
@coroutine_test
|
|
async def test_setting_explicit(self, server_url: str) -> None:
|
|
crawler = get_crawler()
|
|
settings = Settings({"DOWNLOADER_CLIENT_TLS_METHOD": "TLSv1.2"})
|
|
client_context_factory = load_context_factory_from_settings(settings, crawler)
|
|
assert client_context_factory._ssl_method == OpenSSL.SSL.TLSv1_2_METHOD
|
|
await self._assert_factory_works(server_url, client_context_factory)
|
|
|
|
@coroutine_test
|
|
async def test_direct_from_crawler(self, server_url: str) -> None:
|
|
# the setting is ignored
|
|
crawler = get_crawler(settings_dict={"DOWNLOADER_CLIENT_TLS_METHOD": "bad"})
|
|
client_context_factory = build_from_crawler(ScrapyClientContextFactory, crawler)
|
|
assert client_context_factory._ssl_method == OpenSSL.SSL.SSLv23_METHOD
|
|
await self._assert_factory_works(server_url, client_context_factory)
|
|
|
|
@coroutine_test
|
|
async def test_direct_init(self, server_url: str) -> None:
|
|
client_context_factory = ScrapyClientContextFactory(OpenSSL.SSL.TLSv1_2_METHOD)
|
|
assert client_context_factory._ssl_method == OpenSSL.SSL.TLSv1_2_METHOD
|
|
await self._assert_factory_works(server_url, client_context_factory)
|
|
|
|
|
|
@coroutine_test
|
|
async def test_fetch_deprecated_spider_arg():
|
|
class CustomDownloader(Downloader):
|
|
def fetch(self, request, spider): # pylint: disable=signature-differs
|
|
return super().fetch(request, spider)
|
|
|
|
crawler = get_crawler(DefaultSpider, {"DOWNLOADER": CustomDownloader})
|
|
with pytest.warns(
|
|
ScrapyDeprecationWarning,
|
|
match=r"The fetch\(\) method of .+\.CustomDownloader requires a spider argument",
|
|
):
|
|
await crawler.crawl_async()
|