Migrate to pytest-twisted (#6938)

* Migrate to pytest-twisted (WIP)

* Some typing fixes.

* Make --reactor=asyncio the default again.

* Try installing the correct event loop policy in tests on Windows.

* Make reactor_pytest a normal fixture.

* Fix test warnings.

* Fix FTPDownloadHandler teardown.

* Cleanups, typing.

* More cleanup.

* Update only_asyncio/only_not_asyncio mark messages.
This commit is contained in:
Andrey Rakhmatullin 2025-07-06 21:27:17 +05:00 committed by GitHub
parent 14eace5d8f
commit 6b2997af90
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
49 changed files with 1054 additions and 939 deletions

View File

@ -3,7 +3,7 @@ from pathlib import Path
import pytest
from twisted.web.http import H2_ENABLED
from scrapy.utils.reactor import install_reactor
from scrapy.utils.reactor import set_asyncio_event_loop_policy
from tests.keys import generate_keys
@ -48,36 +48,24 @@ if not H2_ENABLED:
)
def pytest_addoption(parser):
parser.addoption(
"--reactor",
default="asyncio",
choices=["default", "asyncio"],
)
@pytest.fixture(scope="class")
def reactor_pytest(request):
if not request.cls:
# doctests
return None
request.cls.reactor_pytest = request.config.getoption("--reactor")
return request.cls.reactor_pytest
@pytest.fixture(scope="session")
def reactor_pytest(request) -> str:
return request.config.getoption("--reactor")
@pytest.fixture(autouse=True)
def only_asyncio(request, reactor_pytest):
if request.node.get_closest_marker("only_asyncio") and reactor_pytest == "default":
pytest.skip("This test is only run without --reactor=default")
if request.node.get_closest_marker("only_asyncio") and reactor_pytest != "asyncio":
pytest.skip("This test is only run with --reactor=asyncio")
@pytest.fixture(autouse=True)
def only_not_asyncio(request, reactor_pytest):
if (
request.node.get_closest_marker("only_not_asyncio")
and reactor_pytest != "default"
and reactor_pytest == "asyncio"
):
pytest.skip("This test is only run with --reactor=default")
pytest.skip("This test is only run without --reactor=asyncio")
@pytest.fixture(autouse=True)
@ -117,11 +105,10 @@ def requires_boto3(request):
def pytest_configure(config):
if config.getoption("--reactor") != "default":
install_reactor("twisted.internet.asyncioreactor.AsyncioSelectorReactor")
else:
# install the reactor explicitly
from twisted.internet import reactor # noqa: F401
if config.getoption("--reactor") == "asyncio":
# Needed on Windows to switch from proactor to selector for Twisted reactor compatibility.
# If we decide to run tests with both, we will need to add a new option and check it here.
set_asyncio_event_loop_policy()
# Generate localhost certificate files, needed by some tests

View File

@ -218,6 +218,9 @@ disable = [
]
[tool.pytest.ini_options]
addopts = [
"--reactor=asyncio",
]
xfail_strict = true
python_files = ["test_*.py", "test_*/__init__.py"]
markers = [

View File

@ -13,7 +13,7 @@ from scrapy.utils.misc import load_object
from scrapy.utils.python import global_object_name
if TYPE_CHECKING:
from asyncio import AbstractEventLoop, AbstractEventLoopPolicy
from asyncio import AbstractEventLoop
from collections.abc import Callable
from twisted.internet.protocol import ServerFactory
@ -100,17 +100,12 @@ def set_asyncio_event_loop_policy() -> None:
so we restrict their use to the absolutely essential case.
This should only be used to install the reactor.
"""
_get_asyncio_event_loop_policy()
def _get_asyncio_event_loop_policy() -> AbstractEventLoopPolicy:
policy = asyncio.get_event_loop_policy()
if sys.platform == "win32" and not isinstance(
policy, asyncio.WindowsSelectorEventLoopPolicy
):
policy = asyncio.WindowsSelectorEventLoopPolicy()
asyncio.set_event_loop_policy(policy)
return policy
def install_reactor(reactor_path: str, event_loop_path: str | None = None) -> None:

View File

@ -3,7 +3,6 @@ from typing import Any
from unittest.mock import patch
from twisted.internet.defer import inlineCallbacks
from twisted.trial import unittest
from scrapy import Spider
from scrapy.crawler import Crawler, CrawlerRunner
@ -52,7 +51,7 @@ class TestAddon:
assert settings["KEY3"] == "addon"
class TestAddonManager(unittest.TestCase):
class TestAddonManager:
def test_load_settings(self):
settings_dict = {
"ADDONS": {"tests.test_addons.SimpleAddon": 0},

View File

@ -1,5 +1,4 @@
from twisted.internet.defer import inlineCallbacks
from twisted.trial.unittest import TestCase
from scrapy.utils.test import get_crawler
from tests.mockserver import MockServer
@ -12,14 +11,14 @@ from tests.spiders import (
)
class TestCloseSpider(TestCase):
class TestCloseSpider:
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
@inlineCallbacks

View File

@ -3,7 +3,6 @@ from unittest import TextTestResult
import pytest
from twisted.internet.defer import inlineCallbacks
from twisted.python import failure
from twisted.trial import unittest
from scrapy import FormRequest
from scrapy.contracts import Contract, ContractsManager
@ -247,7 +246,7 @@ class InheritsDemoSpider(DemoSpider):
name = "inherits_demo_spider"
class TestContractsManager(unittest.TestCase):
class TestContractsManager:
contracts = [
UrlContract,
CallbackKeywordArgumentsContract,
@ -259,7 +258,7 @@ class TestContractsManager(unittest.TestCase):
CustomFailContract,
]
def setUp(self):
def setup_method(self):
self.conman = ContractsManager(self.contracts)
self.results = TextTestResult(stream=None, descriptions=False, verbosity=0)

View File

@ -1,16 +1,12 @@
from __future__ import annotations
import shutil
import warnings
from pathlib import Path
from tempfile import mkdtemp
from typing import TYPE_CHECKING, Any, cast
import OpenSSL.SSL
import pytest
from twisted.internet.defer import Deferred, inlineCallbacks
from pytest_twisted import async_yield_fixture
from twisted.protocols.policies import WrappingFactory
from twisted.trial import unittest
from twisted.web import server, static
from twisted.web.client import Agent, BrowserLikePolicyForHTTPS, readBody
from twisted.web.client import Response as TxResponse
@ -29,6 +25,7 @@ from scrapy.utils.test import get_crawler
from tests.mockserver import PayloadResource, ssl_context_factory
if TYPE_CHECKING:
from twisted.internet.defer import Deferred
from twisted.web.iweb import IBodyProducer
@ -38,9 +35,23 @@ class TestSlot:
assert repr(slot) == "Slot(concurrency=8, delay=0.10, randomize_delay=True)"
class TestContextFactoryBase(unittest.TestCase):
class TestContextFactoryBase:
context_factory = 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)
wrapper = WrappingFactory(site)
port = self._listen(wrapper)
portno = port.getHost().port
yield f"https://127.0.0.1:{portno}/"
await port.stopListening()
def _listen(self, site):
from twisted.internet import reactor
@ -51,24 +62,6 @@ class TestContextFactoryBase(unittest.TestCase):
interface="127.0.0.1",
)
def getURL(self, path):
return f"https://127.0.0.1:{self.portno}/{path}"
def setUp(self):
self.tmpname = Path(mkdtemp())
(self.tmpname / "file").write_bytes(b"0123456789")
r = static.File(str(self.tmpname))
r.putChild(b"payload", PayloadResource())
self.site = server.Site(r, timeout=None)
self.wrapper = WrappingFactory(self.site)
self.port = self._listen(self.wrapper)
self.portno = self.port.getHost().port
@inlineCallbacks
def tearDown(self):
yield self.port.stopListening()
shutil.rmtree(self.tmpname)
@staticmethod
async def get_page(
url: str,
@ -102,13 +95,13 @@ class TestContextFactoryBase(unittest.TestCase):
class TestContextFactory(TestContextFactoryBase):
@deferred_f_from_coro_f
async def testPayload(self):
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(
self.getURL("payload"), client_context_factory, body=s
server_url + "payload", client_context_factory, body=s
)
assert body == to_bytes(s)
@ -131,21 +124,21 @@ class TestContextFactory(TestContextFactoryBase):
class TestContextFactoryTLSMethod(TestContextFactoryBase):
async def _assert_factory_works(
self, client_context_factory: ScrapyClientContextFactory
self, server_url: str, client_context_factory: ScrapyClientContextFactory
) -> None:
s = "0123456789" * 10
body = await self.get_page(
self.getURL("payload"), client_context_factory, body=s
server_url + "payload", client_context_factory, body=s
)
assert body == to_bytes(s)
@deferred_f_from_coro_f
async def test_setting_default(self):
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(client_context_factory)
await self._assert_factory_works(server_url, client_context_factory)
def test_setting_none(self):
crawler = get_crawler()
@ -160,23 +153,23 @@ class TestContextFactoryTLSMethod(TestContextFactoryBase):
load_context_factory_from_settings(settings, crawler)
@deferred_f_from_coro_f
async def test_setting_explicit(self):
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(client_context_factory)
await self._assert_factory_works(server_url, client_context_factory)
@deferred_f_from_coro_f
async def test_direct_from_crawler(self):
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(client_context_factory)
await self._assert_factory_works(server_url, client_context_factory)
@deferred_f_from_coro_f
async def test_direct_init(self):
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(client_context_factory)
await self._assert_factory_works(server_url, client_context_factory)

View File

@ -12,7 +12,6 @@ from testfixtures import LogCapture
from twisted.internet.defer import inlineCallbacks
from twisted.internet.ssl import Certificate
from twisted.python.failure import Failure
from twisted.trial.unittest import TestCase
from scrapy import Spider, signals
from scrapy.crawler import CrawlerRunner
@ -61,16 +60,16 @@ if TYPE_CHECKING:
from scrapy.statscollectors import StatsCollector
class TestCrawl(TestCase):
class TestCrawl:
mockserver: MockServer
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
@inlineCallbacks
@ -422,16 +421,16 @@ with multiples lines
assert "Got response 200" in str(log)
class TestCrawlSpider(TestCase):
class TestCrawlSpider:
mockserver: MockServer
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
async def _run_spider(

View File

@ -7,6 +7,7 @@ import subprocess
import sys
import warnings
from abc import ABC, abstractmethod
from collections.abc import Generator
from pathlib import Path
from typing import Any
@ -14,7 +15,6 @@ import pytest
from packaging.version import parse as parse_version
from pexpect.popen_spawn import PopenSpawn
from twisted.internet.defer import Deferred, inlineCallbacks
from twisted.trial import unittest
from w3lib import __version__ as w3lib_version
from zope.interface.exceptions import MultipleInvalid
@ -48,7 +48,7 @@ def get_raw_crawler(spidercls=None, settings_dict=None):
return Crawler(spidercls or DefaultSpider, settings)
class TestBaseCrawler(unittest.TestCase):
class TestBaseCrawler:
def assertOptionIsDefault(self, settings, key):
assert isinstance(settings, Settings)
assert settings[key] == getattr(default_settings, key)
@ -648,8 +648,7 @@ class NoRequestsSpider(scrapy.Spider):
yield
@pytest.mark.usefixtures("reactor_pytest")
class TestCrawlerRunnerHasSpider(unittest.TestCase):
class TestCrawlerRunnerHasSpider:
@staticmethod
def _runner():
return CrawlerRunner(get_reactor_settings())
@ -700,8 +699,10 @@ class TestCrawlerRunnerHasSpider(unittest.TestCase):
assert runner.bootstrap_failed
@inlineCallbacks
def test_crawler_runner_asyncio_enabled_true(self):
if self.reactor_pytest == "default":
def test_crawler_runner_asyncio_enabled_true(
self, reactor_pytest: str
) -> Generator[Deferred[Any], Any, None]:
if reactor_pytest != "asyncio":
runner = CrawlerRunner(
settings={
"TWISTED_REACTOR": "twisted.internet.asyncioreactor.AsyncioSelectorReactor",
@ -760,7 +761,7 @@ class ScriptRunnerMixin(ABC):
return stderr.decode("utf-8")
class TestCrawlerProcessSubprocessBase(ScriptRunnerMixin, unittest.TestCase):
class TestCrawlerProcessSubprocessBase(ScriptRunnerMixin):
"""Common tests between CrawlerProcess and AsyncCrawlerProcess,
with the same file names and expectations.
"""

View File

@ -8,9 +8,12 @@ import pytest
from scrapy.core.downloader.handlers.http10 import HTTP10DownloadHandler
from scrapy.http import Request
from scrapy.spiders import Spider
from scrapy.utils.defer import deferred_f_from_coro_f
from tests.test_downloader_handlers_http_base import TestHttpBase, TestHttpProxyBase
from tests.test_downloader_handlers_http_base import (
TestHttpBase,
TestHttpProxyBase,
download_request,
)
if TYPE_CHECKING:
from scrapy.core.downloader.handlers import DownloadHandlerProtocol
@ -27,9 +30,11 @@ class TestHttp10(HTTP10DownloadHandlerMixin, TestHttpBase):
"""HTTP 1.0 test case"""
@deferred_f_from_coro_f
async def test_protocol(self):
request = Request(self.getURL("host"), method="GET")
response = await self.download_request(request, Spider("foo"))
async def test_protocol(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "host"), method="GET")
response = await download_request(download_handler, request)
assert response.protocol == "HTTP/1.0"

View File

@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Any
from unittest import mock
import pytest
from pytest_twisted import async_yield_fixture
from testfixtures import LogCapture
from twisted.internet import defer, error
from twisted.web import server
@ -16,8 +17,6 @@ from twisted.web.http import H2_ENABLED
from scrapy.http import Request
from scrapy.spiders import Spider
from scrapy.utils.defer import deferred_f_from_coro_f, maybe_deferred_to_future
from scrapy.utils.misc import build_from_crawler
from scrapy.utils.test import get_crawler
from tests.mockserver import ssl_context_factory
from tests.test_downloader_handlers_http_base import (
TestHttpMockServerBase,
@ -28,9 +27,12 @@ from tests.test_downloader_handlers_http_base import (
TestHttpsInvalidDNSPatternBase,
TestHttpsWrongHostnameBase,
UriResource,
download_request,
)
if TYPE_CHECKING:
from collections.abc import AsyncGenerator
from scrapy.core.downloader.handlers import DownloadHandlerProtocol
@ -54,84 +56,96 @@ class TestHttps2(H2DownloadHandlerMixin, TestHttps11Base):
HTTP2_DATALOSS_SKIP_REASON = "Content-Length mismatch raises InvalidBodyLengthError"
@deferred_f_from_coro_f
async def test_protocol(self):
request = Request(self.getURL("host"), method="GET")
response = await self.download_request(request, Spider("foo"))
async def test_protocol(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "host"), method="GET")
response = await download_request(download_handler, request)
assert response.protocol == "h2"
@deferred_f_from_coro_f
async def test_download_with_maxsize_very_large_file(self):
async def test_download_with_maxsize_very_large_file(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
from twisted.internet import reactor
with mock.patch("scrapy.core.http2.stream.logger") as logger:
request = Request(self.getURL("largechunkedfile"))
request = Request(self.getURL(server_port, "largechunkedfile"))
def check(logger):
def check(logger: mock.Mock) -> None:
logger.error.assert_called_once_with(mock.ANY)
with pytest.raises((defer.CancelledError, error.ConnectionAborted)):
await self.download_request(
request, Spider("foo", download_maxsize=1500)
await download_request(
download_handler, request, Spider("foo", download_maxsize=1500)
)
# As the error message is logged in the dataReceived callback, we
# have to give a bit of time to the reactor to process the queue
# after closing the connection.
d = defer.Deferred()
d: defer.Deferred[mock.Mock] = defer.Deferred()
d.addCallback(check)
reactor.callLater(0.1, d.callback, logger)
await maybe_deferred_to_future(d)
@deferred_f_from_coro_f
async def test_unsupported_scheme(self):
async def test_unsupported_scheme(
self, download_handler: DownloadHandlerProtocol
) -> None:
request = Request("ftp://unsupported.scheme")
with pytest.raises(SchemeNotSupported):
await self.download_request(request, Spider("foo"))
await download_request(download_handler, request)
async def _test_download_cause_data_loss(self, url: str) -> None:
def test_download_cause_data_loss(self) -> None: # type: ignore[override]
pytest.skip(self.HTTP2_DATALOSS_SKIP_REASON)
async def _test_download_allow_data_loss(self, url: str) -> None:
def test_download_allow_data_loss(self) -> None: # type: ignore[override]
pytest.skip(self.HTTP2_DATALOSS_SKIP_REASON)
async def _test_download_allow_data_loss_via_setting(self, url: str) -> None:
def test_download_allow_data_loss_via_setting(self) -> None: # type: ignore[override]
pytest.skip(self.HTTP2_DATALOSS_SKIP_REASON)
@deferred_f_from_coro_f
async def test_concurrent_requests_same_domain(self):
spider = Spider("foo")
request1 = Request(self.getURL("file"))
response1 = await self.download_request(request1, spider)
async def test_concurrent_requests_same_domain(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request1 = Request(self.getURL(server_port, "file"))
response1 = await download_request(download_handler, request1)
assert response1.body == b"0123456789"
request2 = Request(self.getURL("echo"), method="POST")
response2 = await self.download_request(request2, spider)
request2 = Request(self.getURL(server_port, "echo"), method="POST")
response2 = await download_request(download_handler, request2)
assert response2.headers["Content-Length"] == b"79"
@pytest.mark.xfail(reason="https://github.com/python-hyper/h2/issues/1247")
@deferred_f_from_coro_f
async def test_connect_request(self):
request = Request(self.getURL("file"), method="CONNECT")
response = await self.download_request(request, Spider("foo"))
async def test_connect_request(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "file"), method="CONNECT")
response = await download_request(download_handler, request)
assert response.body == b""
@deferred_f_from_coro_f
async def test_custom_content_length_good(self):
request = Request(self.getURL("contentlength"))
async def test_custom_content_length_good(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "contentlength"))
custom_content_length = str(len(request.body))
request.headers["Content-Length"] = custom_content_length
response = await self.download_request(request, Spider("foo"))
response = await download_request(download_handler, request)
assert response.text == custom_content_length
@deferred_f_from_coro_f
async def test_custom_content_length_bad(self):
request = Request(self.getURL("contentlength"))
async def test_custom_content_length_bad(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "contentlength"))
actual_content_length = str(len(request.body))
bad_content_length = str(len(request.body) + 1)
request.headers["Content-Length"] = bad_content_length
with LogCapture() as log:
response = await self.download_request(request, Spider("foo"))
response = await download_request(download_handler, request)
assert response.text == actual_content_length
log.check_present(
(
@ -144,12 +158,14 @@ class TestHttps2(H2DownloadHandlerMixin, TestHttps11Base):
)
@deferred_f_from_coro_f
async def test_duplicate_header(self):
request = Request(self.getURL("echo"))
async def test_duplicate_header(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "echo"))
header, value1, value2 = "Custom-Header", "foo", "bar"
request.headers.appendlist(header, value1)
request.headers.appendlist(header, value2)
response = await self.download_request(request, Spider("foo"))
response = await download_request(download_handler, request)
assert json.loads(response.text)["headers"][header] == [value1, value2]
@ -189,33 +205,32 @@ class TestHttps2Proxy(H2DownloadHandlerMixin, TestHttpProxyBase):
# only used for HTTPS tests
keyfile = "keys/localhost.key"
certfile = "keys/localhost.crt"
scheme = "https"
host = "127.0.0.1"
expected_http_proxy_request_body = b"/"
def setUp(self):
@async_yield_fixture
async def server_port(self) -> AsyncGenerator[int]:
from twisted.internet import reactor
site = server.Site(UriResource(), timeout=None)
self.port = reactor.listenSSL(
port = reactor.listenSSL(
0,
site,
ssl_context_factory(self.keyfile, self.certfile),
interface=self.host,
)
self.portno = self.port.getHost().port
self.download_handler = build_from_crawler(
self.download_handler_cls, get_crawler()
)
def getURL(self, path):
return f"{self.scheme}://{self.host}:{self.portno}/{path}"
yield port.getHost().port
await port.stopListening()
@deferred_f_from_coro_f
async def test_download_with_proxy_https_timeout(self):
async def test_download_with_proxy_https_timeout(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
with pytest.raises(NotImplementedError):
await maybe_deferred_to_future(
super().test_download_with_proxy_https_timeout()
super().test_download_with_proxy_https_timeout(
server_port, download_handler
)
)

View File

@ -4,17 +4,16 @@ from __future__ import annotations
import contextlib
import os
import shutil
import sys
from pathlib import Path
from tempfile import mkdtemp, mkstemp
from typing import TYPE_CHECKING, Any
from unittest import mock
import pytest
from pytest_twisted import async_yield_fixture
from twisted.cred import checkers, credentials, portal
from twisted.internet.defer import inlineCallbacks
from twisted.protocols.ftp import ConnectionLost, FTPFactory, FTPRealm
from twisted.trial import unittest
from w3lib.url import path_to_file_uri
from scrapy.core.downloader.handlers import DownloadHandlers
@ -26,12 +25,15 @@ from scrapy.exceptions import NotConfigured
from scrapy.http import HtmlResponse, Request, Response
from scrapy.http.response.text import TextResponse
from scrapy.responsetypes import responsetypes
from scrapy.spiders import Spider
from scrapy.utils.defer import deferred_f_from_coro_f, 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
if TYPE_CHECKING:
from collections.abc import AsyncGenerator, Generator
class DummyDH:
lazy = False
@ -92,27 +94,27 @@ class TestLoad:
assert "scheme" not in dh._notconfigured
class TestFile(unittest.TestCase):
def setUp(self):
class TestFile:
def setup_method(self):
# add a special char to check that they are handled correctly
self.fd, self.tmpname = mkstemp(suffix="^")
Path(self.tmpname).write_text("0123456789", encoding="utf-8")
self.download_handler = build_from_crawler(FileDownloadHandler, get_crawler())
def tearDown(self):
def teardown_method(self):
os.close(self.fd)
Path(self.tmpname).unlink()
async def download_request(self, request: Request, spider: Spider) -> Response:
async def download_request(self, request: Request) -> Response:
return await maybe_deferred_to_future(
self.download_handler.download_request(request, spider)
self.download_handler.download_request(request, DefaultSpider())
)
@deferred_f_from_coro_f
async def test_download(self):
request = Request(path_to_file_uri(self.tmpname))
assert request.url.upper().endswith("%5E")
response = await self.download_request(request, Spider("foo"))
response = await self.download_request(request)
assert response.url == request.url
assert response.status == 200
assert response.body == b"0123456789"
@ -123,7 +125,7 @@ class TestFile(unittest.TestCase):
request = Request(path_to_file_uri(mkdtemp()))
# the specific exception differs between platforms
with pytest.raises(OSError): # noqa: PT011
await self.download_request(request, Spider("foo"))
await self.download_request(request)
class HttpDownloadHandlerMock:
@ -145,7 +147,7 @@ class TestS3Anon:
# anon=True, # implicit
)
self.download_request = self.s3reqh.download_request
self.spider = Spider("foo")
self.spider = DefaultSpider()
def test_anon_request(self):
req = Request("s3://aws-publicdatasets/")
@ -176,7 +178,7 @@ class TestS3:
httpdownloadhandler=HttpDownloadHandlerMock,
)
self.download_request = s3reqh.download_request
self.spider = Spider("foo")
self.spider = DefaultSpider()
@contextlib.contextmanager
def _mocked_date(self, date):
@ -304,10 +306,10 @@ class TestS3:
)
class TestFTPBase(unittest.TestCase):
class TestFTPBase:
username = "scrapy"
password = "passwd"
req_meta = {"ftp_user": username, "ftp_password": password}
req_meta: dict[str, Any] = {"ftp_user": username, "ftp_password": password}
test_files = (
("file.txt", b"I have the power!"),
@ -315,194 +317,182 @@ class TestFTPBase(unittest.TestCase):
("html-file-without-extension", b"<!DOCTYPE html>\n<title>.</title>"),
)
def setUp(self):
from twisted.internet import reactor
# setup dirs and test file
self.directory = Path(mkdtemp())
userdir = self.directory / self.username
def _create_files(self, root: Path) -> None:
userdir = root / self.username
userdir.mkdir()
for filename, content in self.test_files:
(userdir / filename).write_bytes(content)
# setup server
realm = FTPRealm(
anonymousRoot=str(self.directory), userHome=str(self.directory)
)
def _get_factory(self, root):
realm = FTPRealm(anonymousRoot=str(root), userHome=str(root))
p = portal.Portal(realm)
users_checker = checkers.InMemoryUsernamePasswordDatabaseDontUse()
users_checker.addUser(self.username, self.password)
p.registerChecker(users_checker, credentials.IUsernamePassword)
self.factory = FTPFactory(portal=p)
self.port = reactor.listenTCP(0, self.factory, interface="127.0.0.1")
self.portNum = self.port.getHost().port
return FTPFactory(portal=p)
@async_yield_fixture
async def server_url(self, tmp_path: Path) -> AsyncGenerator[str]:
from twisted.internet import reactor
self._create_files(tmp_path)
factory = self._get_factory(tmp_path)
port = reactor.listenTCP(0, factory, interface="127.0.0.1")
portno = port.getHost().port
yield f"https://127.0.0.1:{portno}/"
await port.stopListening()
@staticmethod
@pytest.fixture
def dh() -> Generator[FTPDownloadHandler]:
crawler = get_crawler()
self.download_handler = build_from_crawler(FTPDownloadHandler, crawler)
dh = build_from_crawler(FTPDownloadHandler, crawler)
@inlineCallbacks
def tearDown(self):
yield self.port.stopListening()
shutil.rmtree(self.directory)
yield dh
async def download_request(self, request: Request) -> Response:
# if the test was skipped, there will be no client attribute
if hasattr(dh, "client"):
assert dh.client.transport
dh.client.transport.loseConnection()
@staticmethod
async def download_request(dh: FTPDownloadHandler, request: Request) -> Response:
return await maybe_deferred_to_future(
self.download_handler.download_request(request, None)
dh.download_request(request, DefaultSpider())
)
def _lose_connection(self):
self.download_handler.client.transport.loseConnection()
@deferred_f_from_coro_f
async def test_ftp_download_success(
self, server_url: str, dh: FTPDownloadHandler
) -> None:
request = Request(url=server_url + "file.txt", meta=self.req_meta)
r = await self.download_request(dh, request)
assert r.status == 200
assert r.body == b"I have the power!"
assert r.headers == {b"Local Filename": [b""], b"Size": [b"17"]}
assert r.protocol is None
@deferred_f_from_coro_f
async def test_ftp_download_success(self):
async def test_ftp_download_path_with_spaces(
self, server_url: str, dh: FTPDownloadHandler
) -> None:
request = Request(
url=f"ftp://127.0.0.1:{self.portNum}/file.txt", meta=self.req_meta
)
try:
r = await self.download_request(request)
assert r.status == 200
assert r.body == b"I have the power!"
assert r.headers == {b"Local Filename": [b""], b"Size": [b"17"]}
assert r.protocol is None
finally:
self._lose_connection()
@deferred_f_from_coro_f
async def test_ftp_download_path_with_spaces(self):
request = Request(
url=f"ftp://127.0.0.1:{self.portNum}/file with spaces.txt",
url=server_url + "file with spaces.txt",
meta=self.req_meta,
)
try:
r = await self.download_request(request)
assert r.status == 200
assert r.body == b"Moooooooooo power!"
assert r.headers == {b"Local Filename": [b""], b"Size": [b"18"]}
finally:
self._lose_connection()
r = await self.download_request(dh, request)
assert r.status == 200
assert r.body == b"Moooooooooo power!"
assert r.headers == {b"Local Filename": [b""], b"Size": [b"18"]}
@deferred_f_from_coro_f
async def test_ftp_download_nonexistent(self):
request = Request(
url=f"ftp://127.0.0.1:{self.portNum}/nonexistent.txt", meta=self.req_meta
)
try:
r = await self.download_request(request)
assert r.status == 404
finally:
self._lose_connection()
async def test_ftp_download_nonexistent(
self, server_url: str, dh: FTPDownloadHandler
) -> None:
request = Request(url=server_url + "nonexistent.txt", meta=self.req_meta)
r = await self.download_request(dh, request)
assert r.status == 404
@deferred_f_from_coro_f
async def test_ftp_local_filename(self):
async def test_ftp_local_filename(
self, server_url: str, dh: FTPDownloadHandler
) -> None:
f, local_fname = mkstemp()
fname_bytes = to_bytes(local_fname)
local_fname = Path(local_fname)
local_path = Path(local_fname)
os.close(f)
meta = {"ftp_local_filename": fname_bytes}
meta.update(self.req_meta)
request = Request(url=f"ftp://127.0.0.1:{self.portNum}/file.txt", meta=meta)
try:
r = await self.download_request(request)
assert r.body == fname_bytes
assert r.headers == {b"Local Filename": [fname_bytes], b"Size": [b"17"]}
assert local_fname.exists()
assert local_fname.read_bytes() == b"I have the power!"
local_fname.unlink()
finally:
self._lose_connection()
request = Request(url=server_url + "file.txt", meta=meta)
r = await self.download_request(dh, request)
assert r.body == fname_bytes
assert r.headers == {b"Local Filename": [fname_bytes], b"Size": [b"17"]}
assert local_path.exists()
assert local_path.read_bytes() == b"I have the power!"
local_path.unlink()
async def _test_response_class(self, filename: str, response_class: type[Response]):
@pytest.mark.parametrize(
("filename", "response_class"),
[
("file.txt", TextResponse),
("html-file-without-extension", HtmlResponse),
],
)
@deferred_f_from_coro_f
async def test_response_class(
self,
filename: str,
response_class: type[Response],
server_url: str,
dh: FTPDownloadHandler,
) -> None:
f, local_fname = mkstemp()
local_fname_path = Path(local_fname)
os.close(f)
meta = {}
meta.update(self.req_meta)
request = Request(url=f"ftp://127.0.0.1:{self.portNum}/{filename}", meta=meta)
try:
r = await self.download_request(request)
assert type(r) is response_class # pylint: disable=unidiomatic-typecheck
local_fname_path.unlink()
finally:
self._lose_connection()
@deferred_f_from_coro_f
async def test_response_class_from_url(self):
await self._test_response_class("file.txt", TextResponse)
@deferred_f_from_coro_f
async def test_response_class_from_body(self):
await self._test_response_class("html-file-without-extension", HtmlResponse)
request = Request(url=server_url + filename, meta=meta)
r = await self.download_request(dh, request)
assert type(r) is response_class # pylint: disable=unidiomatic-typecheck
local_fname_path.unlink()
class TestFTP(TestFTPBase):
@deferred_f_from_coro_f
async def test_invalid_credentials(self):
if self.reactor_pytest != "default" and sys.platform == "win32":
async def test_invalid_credentials(
self, server_url: str, dh: FTPDownloadHandler, reactor_pytest: str
) -> None:
if reactor_pytest == "asyncio" and sys.platform == "win32":
pytest.skip(
"This test produces DirtyReactorAggregateError on Windows with asyncio"
)
meta = dict(self.req_meta)
meta.update({"ftp_password": "invalid"})
request = Request(url=f"ftp://127.0.0.1:{self.portNum}/file.txt", meta=meta)
try:
with pytest.raises(ConnectionLost):
await self.download_request(request)
finally:
self._lose_connection()
request = Request(url=server_url + "file.txt", meta=meta)
with pytest.raises(ConnectionLost):
await self.download_request(dh, request)
class TestAnonymousFTP(TestFTPBase):
username = "anonymous"
req_meta = {}
def setUp(self):
from twisted.internet import reactor
# setup dir and test file
self.directory = Path(mkdtemp())
def _create_files(self, root: Path) -> None:
for filename, content in self.test_files:
(self.directory / filename).write_bytes(content)
(root / filename).write_bytes(content)
# setup server for anonymous access
realm = FTPRealm(anonymousRoot=str(self.directory))
def _get_factory(self, tmp_path):
realm = FTPRealm(anonymousRoot=str(tmp_path))
p = portal.Portal(realm)
p.registerChecker(checkers.AllowAnonymousAccess(), credentials.IAnonymous)
self.factory = FTPFactory(portal=p, userAnonymous=self.username)
self.port = reactor.listenTCP(0, self.factory, interface="127.0.0.1")
self.portNum = self.port.getHost().port
crawler = get_crawler()
self.download_handler = build_from_crawler(FTPDownloadHandler, crawler)
@inlineCallbacks
def tearDown(self):
yield self.port.stopListening()
shutil.rmtree(self.directory)
return FTPFactory(portal=p, userAnonymous=self.username)
class TestDataURI(unittest.TestCase):
def setUp(self):
class TestDataURI:
def setup_method(self):
crawler = get_crawler()
self.download_handler = build_from_crawler(DataURIDownloadHandler, crawler)
self.spider = Spider("foo")
async def download_request(self, request: Request, spider: Spider) -> Response:
async def download_request(self, request: Request) -> Response:
return await maybe_deferred_to_future(
self.download_handler.download_request(request, spider)
self.download_handler.download_request(request, DefaultSpider())
)
@deferred_f_from_coro_f
async def test_response_attrs(self):
uri = "data:,A%20brief%20note"
request = Request(uri)
response = await self.download_request(request, self.spider)
response = await self.download_request(request)
assert response.url == uri
assert not response.headers
@deferred_f_from_coro_f
async def test_default_mediatype_encoding(self):
request = Request("data:,A%20brief%20note")
response = await self.download_request(request, self.spider)
response = await self.download_request(request)
assert response.text == "A brief note"
assert type(response) is responsetypes.from_mimetype("text/plain") # pylint: disable=unidiomatic-typecheck
assert response.encoding == "US-ASCII"
@ -510,7 +500,7 @@ class TestDataURI(unittest.TestCase):
@deferred_f_from_coro_f
async def test_default_mediatype(self):
request = Request("data:;charset=iso-8859-7,%be%d3%be")
response = await self.download_request(request, self.spider)
response = await self.download_request(request)
assert response.text == "\u038e\u03a3\u038e"
assert type(response) is responsetypes.from_mimetype("text/plain") # pylint: disable=unidiomatic-typecheck
assert response.encoding == "iso-8859-7"
@ -518,7 +508,7 @@ class TestDataURI(unittest.TestCase):
@deferred_f_from_coro_f
async def test_text_charset(self):
request = Request("data:text/plain;charset=iso-8859-7,%be%d3%be")
response = await self.download_request(request, self.spider)
response = await self.download_request(request)
assert response.text == "\u038e\u03a3\u038e"
assert response.body == b"\xbe\xd3\xbe"
assert response.encoding == "iso-8859-7"
@ -530,7 +520,7 @@ class TestDataURI(unittest.TestCase):
"charset=utf-8;bar=%22foo;%5C%22 foo ;/,%22"
",%CE%8E%CE%A3%CE%8E"
)
response = await self.download_request(request, self.spider)
response = await self.download_request(request)
assert response.text == "\u038e\u03a3\u038e"
assert type(response) is responsetypes.from_mimetype("text/plain") # pylint: disable=unidiomatic-typecheck
assert response.encoding == "utf-8"
@ -538,11 +528,11 @@ class TestDataURI(unittest.TestCase):
@deferred_f_from_coro_f
async def test_base64(self):
request = Request("data:text/plain;base64,SGVsbG8sIHdvcmxkLg%3D%3D")
response = await self.download_request(request, self.spider)
response = await self.download_request(request)
assert response.text == "Hello, world."
@deferred_f_from_coro_f
async def test_protocol(self):
request = Request("data:,")
response = await self.download_request(request, self.spider)
response = await self.download_request(request)
assert response.protocol is None

View File

@ -3,20 +3,16 @@
from __future__ import annotations
import json
import shutil
import sys
from abc import ABC, abstractmethod
from pathlib import Path
from tempfile import mkdtemp
from typing import TYPE_CHECKING, Any
from unittest import mock
import pytest
from pytest_twisted import async_yield_fixture
from testfixtures import LogCapture
from twisted.internet import defer, error
from twisted.internet.defer import inlineCallbacks, maybeDeferred
from twisted.protocols.policies import WrappingFactory
from twisted.trial import unittest
from twisted.web import resource, server, static, util
from twisted.web._newclient import ResponseFailed
from twisted.web.http import _DataLoss
@ -30,6 +26,7 @@ from scrapy.utils.defer import (
)
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 import NON_EXISTING_RESOLVABLE
from tests.mockserver import (
@ -44,6 +41,9 @@ from tests.mockserver import (
from tests.spiders import SingleRequestSpider
if TYPE_CHECKING:
from collections.abc import AsyncGenerator
from pathlib import Path
from scrapy.core.downloader.handlers import DownloadHandlerProtocol
@ -135,8 +135,30 @@ class DuplicateHeaderResource(resource.Resource):
return b""
class TestHttpBase(unittest.TestCase, ABC):
async def download_request(
download_handler: DownloadHandlerProtocol,
request: Request,
spider: Spider = DefaultSpider(),
) -> Response:
return await maybe_deferred_to_future(
download_handler.download_request(request, spider)
)
async def close_dh(dh: DownloadHandlerProtocol) -> None:
# needed because the interface of close() is not clearly defined
if not hasattr(dh, "close"):
return
c = dh.close()
if c is None:
return
# covers coroutines and Deferreds; won't work if close() uses Futures inside
await c
class TestHttpBase(ABC):
scheme = "http"
host = "localhost"
# only used for HTTPS tests
keyfile = "keys/localhost.key"
@ -147,12 +169,10 @@ class TestHttpBase(unittest.TestCase, ABC):
def download_handler_cls(self) -> type[DownloadHandlerProtocol]:
raise NotImplementedError
def setUp(self):
from twisted.internet import reactor
self.tmpname = Path(mkdtemp())
(self.tmpname / "file").write_bytes(b"0123456789")
r = static.File(str(self.tmpname))
@pytest.fixture
def site(self, tmp_path):
(tmp_path / "file").write_bytes(b"0123456789")
r = static.File(str(tmp_path))
r.putChild(b"redirect", util.Redirect(b"/file"))
r.putChild(b"wait", ForeverTakingResource())
r.putChild(b"hang-after-headers", ForeverTakingResource(write=True))
@ -167,112 +187,134 @@ class TestHttpBase(unittest.TestCase, ABC):
r.putChild(b"largechunkedfile", LargeChunkedFileResource())
r.putChild(b"duplicate-header", DuplicateHeaderResource())
r.putChild(b"echo", Echo())
self.site = server.Site(r, timeout=None)
self.wrapper = WrappingFactory(self.site)
self.host = "localhost"
return server.Site(r, timeout=None)
@async_yield_fixture
async def server_port(self, site: server.Site) -> AsyncGenerator[int]:
from twisted.internet import reactor
if self.scheme == "https":
# Using WrappingFactory do not enable HTTP/2 failing all the
# tests with H2DownloadHandler
self.port = reactor.listenSSL(
port = reactor.listenSSL(
0,
self.site,
site,
ssl_context_factory(self.keyfile, self.certfile),
interface=self.host,
)
else:
self.port = reactor.listenTCP(0, self.wrapper, interface=self.host)
self.portno = self.port.getHost().port
self.download_handler = build_from_crawler(
self.download_handler_cls, get_crawler()
)
wrapper = WrappingFactory(site)
port = reactor.listenTCP(0, wrapper, interface=self.host)
@inlineCallbacks
def tearDown(self):
yield self.port.stopListening()
if hasattr(self.download_handler, "close"):
yield self.download_handler.close()
shutil.rmtree(self.tmpname)
yield port.getHost().port
def getURL(self, path):
return f"{self.scheme}://{self.host}:{self.portno}/{path}"
await port.stopListening()
async def download_request(self, request: Request, spider: Spider) -> Response:
return await maybe_deferred_to_future(
self.download_handler.download_request(request, spider)
)
@async_yield_fixture
async def download_handler(self) -> AsyncGenerator[DownloadHandlerProtocol]:
dh = build_from_crawler(self.download_handler_cls, get_crawler())
yield dh
await close_dh(dh)
def getURL(self, portno: int, path: str) -> str:
return f"{self.scheme}://{self.host}:{portno}/{path}"
@deferred_f_from_coro_f
async def test_download(self):
request = Request(self.getURL("file"))
response = await self.download_request(request, Spider("foo"))
async def test_download(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "file"))
response = await download_request(download_handler, request)
assert response.body == b"0123456789"
@deferred_f_from_coro_f
async def test_download_head(self):
request = Request(self.getURL("file"), method="HEAD")
response = await self.download_request(request, Spider("foo"))
async def test_download_head(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "file"), method="HEAD")
response = await download_request(download_handler, request)
assert response.body == b""
@deferred_f_from_coro_f
async def test_redirect_status(self):
request = Request(self.getURL("redirect"))
response = await self.download_request(request, Spider("foo"))
async def test_redirect_status(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "redirect"))
response = await download_request(download_handler, request)
assert response.status == 302
@deferred_f_from_coro_f
async def test_redirect_status_head(self):
request = Request(self.getURL("redirect"), method="HEAD")
response = await self.download_request(request, Spider("foo"))
async def test_redirect_status_head(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "redirect"), method="HEAD")
response = await download_request(download_handler, request)
assert response.status == 302
@deferred_f_from_coro_f
async def test_timeout_download_from_spider_nodata_rcvd(self):
if self.reactor_pytest != "default" and sys.platform == "win32":
async def test_timeout_download_from_spider_nodata_rcvd(
self,
server_port: int,
download_handler: DownloadHandlerProtocol,
reactor_pytest: str,
) -> None:
if reactor_pytest == "asyncio" and sys.platform == "win32":
# https://twistedmatrix.com/trac/ticket/10279
pytest.skip(
"This test produces DirtyReactorAggregateError on Windows with asyncio"
)
# client connects but no data is received
spider = Spider("foo")
meta = {"download_timeout": 0.5}
request = Request(self.getURL("wait"), meta=meta)
d = deferred_from_coro(self.download_request(request, spider))
request = Request(self.getURL(server_port, "wait"), meta=meta)
d = deferred_from_coro(download_request(download_handler, request))
with pytest.raises((defer.TimeoutError, error.TimeoutError)):
await maybe_deferred_to_future(d)
@deferred_f_from_coro_f
async def test_timeout_download_from_spider_server_hangs(self):
if self.reactor_pytest != "default" and sys.platform == "win32":
async def test_timeout_download_from_spider_server_hangs(
self,
server_port: int,
download_handler: DownloadHandlerProtocol,
reactor_pytest: str,
) -> None:
if reactor_pytest == "asyncio" and sys.platform == "win32":
# https://twistedmatrix.com/trac/ticket/10279
pytest.skip(
"This test produces DirtyReactorAggregateError on Windows with asyncio"
)
# client connects, server send headers and some body bytes but hangs
spider = Spider("foo")
meta = {"download_timeout": 0.5}
request = Request(self.getURL("hang-after-headers"), meta=meta)
d = deferred_from_coro(self.download_request(request, spider))
request = Request(self.getURL(server_port, "hang-after-headers"), meta=meta)
d = deferred_from_coro(download_request(download_handler, request))
with pytest.raises((defer.TimeoutError, error.TimeoutError)):
await maybe_deferred_to_future(d)
@deferred_f_from_coro_f
async def test_host_header_not_in_request_headers(self):
request = Request(self.getURL("host"))
response = await self.download_request(request, Spider("foo"))
assert response.body == to_bytes(f"{self.host}:{self.portno}")
async def test_host_header_not_in_request_headers(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "host"))
response = await download_request(download_handler, request)
assert response.body == to_bytes(f"{self.host}:{server_port}")
assert not request.headers
@deferred_f_from_coro_f
async def test_host_header_set_in_request_headers(self):
host = self.host + ":" + str(self.portno)
request = Request(self.getURL("host"), headers={"Host": host})
response = await self.download_request(request, Spider("foo"))
async def test_host_header_set_in_request_headers(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
host = f"{self.host}:{server_port}"
request = Request(self.getURL(server_port, "host"), headers={"Host": host})
response = await download_request(download_handler, request)
assert response.body == host.encode()
assert request.headers.get("Host") == host.encode()
@deferred_f_from_coro_f
async def test_content_length_zero_bodyless_post_request_headers(self):
async def test_content_length_zero_bodyless_post_request_headers(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
"""Tests if "Content-Length: 0" is sent for bodyless POST requests.
This is not strictly required by HTTP RFCs but can cause trouble
@ -283,55 +325,64 @@ class TestHttpBase(unittest.TestCase, ABC):
https://github.com/kennethreitz/requests/issues/405
https://bugs.python.org/issue14721
"""
request = Request(self.getURL("contentlength"), method="POST")
response = await self.download_request(request, Spider("foo"))
request = Request(self.getURL(server_port, "contentlength"), method="POST")
response = await download_request(download_handler, request)
assert response.body == b"0"
@deferred_f_from_coro_f
async def test_content_length_zero_bodyless_post_only_one(self):
request = Request(self.getURL("echo"), method="POST")
response = await self.download_request(request, Spider("foo"))
async def test_content_length_zero_bodyless_post_only_one(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "echo"), method="POST")
response = await download_request(download_handler, request)
headers = Headers(json.loads(response.text)["headers"])
contentlengths = headers.getlist("Content-Length")
assert len(contentlengths) == 1
assert contentlengths == [b"0"]
@deferred_f_from_coro_f
async def test_payload(self):
async def test_payload(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
body = b"1" * 100 # PayloadResource requires body length to be 100
request = Request(self.getURL("payload"), method="POST", body=body)
response = await self.download_request(request, Spider("foo"))
request = Request(self.getURL(server_port, "payload"), method="POST", body=body)
response = await download_request(download_handler, request)
assert response.body == body
@deferred_f_from_coro_f
async def test_response_header_content_length(self):
request = Request(self.getURL("file"), method=b"GET")
response = await self.download_request(request, Spider("foo"))
assert response.headers[b"content-length"] == b"159"
async def _test_response_class(
self, filename: str, body: bytes, response_class: type[Response]
async def test_response_header_content_length(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(filename), body=body)
response = await self.download_request(request, Spider("foo"))
request = Request(self.getURL(server_port, "file"), method="GET")
response = await download_request(download_handler, request)
assert response.headers[b"content-length"] == b"10"
@pytest.mark.parametrize(
("filename", "body", "response_class"),
[
("foo.html", b"", HtmlResponse),
("foo", b"<!DOCTYPE html>\n<title>.</title>", HtmlResponse),
],
)
@deferred_f_from_coro_f
async def test_response_class(
self,
filename: str,
body: bytes,
response_class: type[Response],
server_port: int,
download_handler: DownloadHandlerProtocol,
) -> None:
request = Request(self.getURL(server_port, filename), body=body)
response = await download_request(download_handler, request)
assert type(response) is response_class # pylint: disable=unidiomatic-typecheck
@deferred_f_from_coro_f
async def test_response_class_from_url(self):
await self._test_response_class("foo.html", b"", HtmlResponse)
@deferred_f_from_coro_f
async def test_response_class_from_body(self):
await self._test_response_class(
"foo",
b"<!DOCTYPE html>\n<title>.</title>",
HtmlResponse,
)
@deferred_f_from_coro_f
async def test_get_duplicate_header(self):
request = Request(self.getURL("duplicate-header"))
response = await self.download_request(request, Spider("foo"))
async def test_get_duplicate_header(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "duplicate-header"))
response = await download_request(download_handler, request)
assert response.headers.getlist(b"Set-Cookie") == [b"a=b", b"c=d"]
@ -339,135 +390,152 @@ class TestHttp11Base(TestHttpBase):
"""HTTP 1.1 test case"""
@deferred_f_from_coro_f
async def test_download_without_maxsize_limit(self):
request = Request(self.getURL("file"))
response = await self.download_request(request, Spider("foo"))
async def test_download_without_maxsize_limit(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "file"))
response = await download_request(download_handler, request)
assert response.body == b"0123456789"
@deferred_f_from_coro_f
async def test_response_class_choosing_request(self):
async def test_response_class_choosing_request(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
"""Tests choosing of correct response type
in case of Content-Type is empty but body contains text.
"""
body = b"Some plain text\ndata with tabs\t and null bytes\0"
request = Request(self.getURL("nocontenttype"), body=body)
response = await self.download_request(request, Spider("foo"))
request = Request(self.getURL(server_port, "nocontenttype"), body=body)
response = await download_request(download_handler, request)
assert type(response) is TextResponse # pylint: disable=unidiomatic-typecheck
@deferred_f_from_coro_f
async def test_download_with_maxsize(self):
request = Request(self.getURL("file"))
async def test_download_with_maxsize(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "file"))
# 10 is minimal size for this request and the limit is only counted on
# response body. (regardless of headers)
response = await self.download_request(
request, Spider("foo", download_maxsize=10)
response = await download_request(
download_handler, request, Spider("foo", download_maxsize=10)
)
assert response.body == b"0123456789"
with pytest.raises((defer.CancelledError, error.ConnectionAborted)):
await self.download_request(request, Spider("foo", download_maxsize=9))
await download_request(
download_handler, request, Spider("foo", download_maxsize=9)
)
@deferred_f_from_coro_f
async def test_download_with_maxsize_very_large_file(self):
async def test_download_with_maxsize_very_large_file(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
from twisted.internet import reactor
# TODO: the logger check is specific to scrapy.core.downloader.handlers.http11
with mock.patch("scrapy.core.downloader.handlers.http11.logger") as logger:
request = Request(self.getURL("largechunkedfile"))
request = Request(self.getURL(server_port, "largechunkedfile"))
def check(logger):
def check(logger: mock.Mock) -> None:
logger.warning.assert_called_once_with(mock.ANY, mock.ANY)
with pytest.raises((defer.CancelledError, error.ConnectionAborted)):
await self.download_request(
request, Spider("foo", download_maxsize=1500)
await download_request(
download_handler, request, Spider("foo", download_maxsize=1500)
)
# As the error message is logged in the dataReceived callback, we
# have to give a bit of time to the reactor to process the queue
# after closing the connection.
d = defer.Deferred()
d: defer.Deferred[mock.Mock] = defer.Deferred()
d.addCallback(check)
reactor.callLater(0.1, d.callback, logger)
await maybe_deferred_to_future(d)
@deferred_f_from_coro_f
async def test_download_with_maxsize_per_req(self):
async def test_download_with_maxsize_per_req(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
meta = {"download_maxsize": 2}
request = Request(self.getURL("file"), meta=meta)
request = Request(self.getURL(server_port, "file"), meta=meta)
with pytest.raises((defer.CancelledError, error.ConnectionAborted)):
await self.download_request(request, Spider("foo"))
await download_request(download_handler, request)
@deferred_f_from_coro_f
async def test_download_with_small_maxsize_per_spider(self):
request = Request(self.getURL("file"))
async def test_download_with_small_maxsize_per_spider(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "file"))
with pytest.raises((defer.CancelledError, error.ConnectionAborted)):
await self.download_request(request, Spider("foo", download_maxsize=2))
await download_request(
download_handler, request, Spider("foo", download_maxsize=2)
)
@deferred_f_from_coro_f
async def test_download_with_large_maxsize_per_spider(self):
request = Request(self.getURL("file"))
response = await self.download_request(
request, Spider("foo", download_maxsize=100)
async def test_download_with_large_maxsize_per_spider(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "file"))
response = await download_request(
download_handler, request, Spider("foo", download_maxsize=100)
)
assert response.body == b"0123456789"
@deferred_f_from_coro_f
async def test_download_chunked_content(self):
request = Request(self.getURL("chunked"))
response = await self.download_request(request, Spider("foo"))
async def test_download_chunked_content(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "chunked"))
response = await download_request(download_handler, request)
assert response.body == b"chunked content\n"
async def _test_download_cause_data_loss(self, url: str) -> None:
@pytest.mark.parametrize("url", ["broken", "broken-chunked"])
@deferred_f_from_coro_f
async def test_download_cause_data_loss(
self, url: str, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
# TODO: this one checks for Twisted-specific exceptions
request = Request(self.getURL(url))
request = Request(self.getURL(server_port, url))
with pytest.raises(ResponseFailed) as exc_info:
await self.download_request(request, Spider("foo"))
await download_request(download_handler, request)
assert any(r.check(_DataLoss) for r in exc_info.value.reasons)
@pytest.mark.parametrize("url", ["broken", "broken-chunked"])
@deferred_f_from_coro_f
async def test_download_broken_content_cause_data_loss(self) -> None:
await self._test_download_cause_data_loss("broken")
@deferred_f_from_coro_f
async def test_download_broken_chunked_content_cause_data_loss(self):
await self._test_download_cause_data_loss("broken-chunked")
async def _test_download_allow_data_loss(self, url: str) -> None:
request = Request(self.getURL(url), meta={"download_fail_on_dataloss": False})
response = await self.download_request(request, Spider("foo"))
async def test_download_allow_data_loss(
self, url: str, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(
self.getURL(server_port, url), meta={"download_fail_on_dataloss": False}
)
response = await download_request(download_handler, request)
assert response.flags == ["dataloss"]
@pytest.mark.parametrize("url", ["broken", "broken-chunked"])
@deferred_f_from_coro_f
async def test_download_broken_content_allow_data_loss(self) -> None:
await self._test_download_allow_data_loss("broken")
@deferred_f_from_coro_f
async def test_download_broken_chunked_content_allow_data_loss(self):
await self._test_download_allow_data_loss("broken-chunked")
async def _test_download_allow_data_loss_via_setting(self, url: str) -> None:
async def test_download_allow_data_loss_via_setting(
self, url: str, server_port: int
) -> None:
crawler = get_crawler(settings_dict={"DOWNLOAD_FAIL_ON_DATALOSS": False})
download_handler = build_from_crawler(self.download_handler_cls, crawler)
request = Request(self.getURL(url))
response = await maybe_deferred_to_future(
download_handler.download_request(request, Spider("foo"))
)
request = Request(self.getURL(server_port, url))
try:
response = await maybe_deferred_to_future(
download_handler.download_request(request, DefaultSpider())
)
finally:
d = download_handler.close() # type: ignore[attr-defined]
if d is not None:
await maybe_deferred_to_future(d)
assert response.flags == ["dataloss"]
@deferred_f_from_coro_f
async def test_download_broken_content_allow_data_loss_via_setting(self):
await self._test_download_allow_data_loss_via_setting("broken-chunked")
@deferred_f_from_coro_f
async def test_download_broken_chunked_content_allow_data_loss_via_setting(self):
await self._test_download_allow_data_loss_via_setting("broken-chunked")
@deferred_f_from_coro_f
async def test_protocol(self):
request = Request(self.getURL("host"), method="GET")
response = await self.download_request(request, Spider("foo"))
async def test_protocol(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "host"), method="GET")
response = await download_request(download_handler, request)
assert response.protocol == "HTTP/1.1"
@ -480,30 +548,33 @@ class TestHttps11Base(TestHttp11Base):
)
@deferred_f_from_coro_f
async def test_tls_logging(self):
async def test_tls_logging(self, server_port: int) -> None:
crawler = get_crawler(
settings_dict={"DOWNLOADER_CLIENT_TLS_VERBOSE_LOGGING": True}
)
download_handler = build_from_crawler(self.download_handler_cls, crawler)
try:
with LogCapture() as log_capture:
request = Request(self.getURL("file"))
request = Request(self.getURL(server_port, "file"))
response = await maybe_deferred_to_future(
download_handler.download_request(request, Spider("foo"))
download_handler.download_request(request, DefaultSpider())
)
assert response.body == b"0123456789"
log_capture.check_present(
("scrapy.core.downloader.tls", "DEBUG", self.tls_log_message)
)
finally:
await maybe_deferred_to_future(maybeDeferred(download_handler.close))
d = download_handler.close() # type: ignore[attr-defined]
if d is not None:
await maybe_deferred_to_future(d)
class TestSimpleHttpsBase(unittest.TestCase, ABC):
class TestSimpleHttpsBase(ABC):
"""Base class for special cases tested with just one simple request"""
keyfile = "keys/localhost.key"
certfile = "keys/localhost.crt"
host = "localhost"
cipher_string: str | None = None
@property
@ -511,49 +582,48 @@ class TestSimpleHttpsBase(unittest.TestCase, ABC):
def download_handler_cls(self) -> type[DownloadHandlerProtocol]:
raise NotImplementedError
def setUp(self):
@async_yield_fixture
async def server_port(self, tmp_path: Path) -> AsyncGenerator[int]:
from twisted.internet import reactor
self.tmpname = Path(mkdtemp())
(self.tmpname / "file").write_bytes(b"0123456789")
r = static.File(str(self.tmpname))
self.site = server.Site(r, timeout=None)
self.host = "localhost"
self.port = reactor.listenSSL(
(tmp_path / "file").write_bytes(b"0123456789")
r = static.File(str(tmp_path))
site = server.Site(r, timeout=None)
port = reactor.listenSSL(
0,
self.site,
site,
ssl_context_factory(
self.keyfile, self.certfile, cipher_string=self.cipher_string
),
interface=self.host,
)
self.portno = self.port.getHost().port
yield port.getHost().port
await port.stopListening()
@async_yield_fixture
async def download_handler(self) -> AsyncGenerator[DownloadHandlerProtocol]:
if self.cipher_string is not None:
settings_dict = {"DOWNLOADER_CLIENT_TLS_CIPHERS": self.cipher_string}
else:
settings_dict = None
crawler = get_crawler(settings_dict=settings_dict)
self.download_handler = build_from_crawler(self.download_handler_cls, crawler)
dh = build_from_crawler(self.download_handler_cls, crawler)
@inlineCallbacks
def tearDown(self):
yield self.port.stopListening()
if hasattr(self.download_handler, "close"):
yield self.download_handler.close()
shutil.rmtree(self.tmpname)
yield dh
def getURL(self, path):
return f"https://{self.host}:{self.portno}/{path}"
await close_dh(dh)
async def download_request(self, request: Request, spider: Spider) -> Response:
return await maybe_deferred_to_future(
self.download_handler.download_request(request, spider)
)
def getURL(self, portno: int, path: str) -> str:
return f"https://{self.host}:{portno}/{path}"
@deferred_f_from_coro_f
async def test_download(self):
request = Request(self.getURL("file"))
response = await self.download_request(request, Spider("foo"))
async def test_download(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "file"))
response = await download_request(download_handler, request)
assert response.body == b"0123456789"
@ -570,9 +640,7 @@ class TestHttpsWrongHostnameBase(TestSimpleHttpsBase):
class TestHttpsInvalidDNSIdBase(TestSimpleHttpsBase):
"""Connect to HTTPS hosts with IP while certificate uses domain names IDs."""
def setUp(self):
super().setUp()
self.host = "127.0.0.1"
host = "127.0.0.1"
class TestHttpsInvalidDNSPatternBase(TestSimpleHttpsBase):
@ -586,7 +654,7 @@ class TestHttpsCustomCiphersBase(TestSimpleHttpsBase):
cipher_string = "CAMELLIA256-SHA"
class TestHttpMockServerBase(unittest.TestCase, ABC):
class TestHttpMockServerBase(ABC):
"""HTTP 1.1 test case with MockServer"""
@property
@ -597,12 +665,12 @@ class TestHttpMockServerBase(unittest.TestCase, ABC):
is_secure = False
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
@deferred_f_from_coro_f
@ -650,7 +718,9 @@ class UriResource(resource.Resource):
return b""
class TestHttpProxyBase(unittest.TestCase, ABC):
class TestHttpProxyBase(ABC):
scheme = "http"
host = "127.0.0.1"
expected_http_proxy_request_body = b"http://example.com"
@property
@ -658,64 +728,70 @@ class TestHttpProxyBase(unittest.TestCase, ABC):
def download_handler_cls(self) -> type[DownloadHandlerProtocol]:
raise NotImplementedError
def setUp(self):
@async_yield_fixture
async def server_port(self) -> AsyncGenerator[int]:
from twisted.internet import reactor
site = server.Site(UriResource(), timeout=None)
wrapper = WrappingFactory(site)
self.port = reactor.listenTCP(0, wrapper, interface="127.0.0.1")
self.portno = self.port.getHost().port
self.download_handler = build_from_crawler(
self.download_handler_cls, get_crawler()
)
port = reactor.listenTCP(0, wrapper, interface=self.host)
@inlineCallbacks
def tearDown(self):
yield self.port.stopListening()
if hasattr(self.download_handler, "close"):
yield self.download_handler.close()
yield port.getHost().port
def getURL(self, path):
return f"http://127.0.0.1:{self.portno}/{path}"
await port.stopListening()
async def download_request(self, request: Request, spider: Spider) -> Response:
return await maybe_deferred_to_future(
self.download_handler.download_request(request, spider)
)
@async_yield_fixture
async def download_handler(self) -> AsyncGenerator[DownloadHandlerProtocol]:
dh = build_from_crawler(self.download_handler_cls, get_crawler())
yield dh
await close_dh(dh)
def getURL(self, portno: int, path: str) -> str:
return f"{self.scheme}://{self.host}:{portno}/{path}"
@deferred_f_from_coro_f
async def test_download_with_proxy(self):
http_proxy = self.getURL("")
async def test_download_with_proxy(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
http_proxy = self.getURL(server_port, "")
request = Request("http://example.com", meta={"proxy": http_proxy})
response = await self.download_request(request, Spider("foo"))
response = await download_request(download_handler, request)
assert response.status == 200
assert response.url == request.url
assert response.body == self.expected_http_proxy_request_body
@deferred_f_from_coro_f
async def test_download_without_proxy(self):
request = Request(self.getURL("path/to/resource"))
response = await self.download_request(request, Spider("foo"))
async def test_download_without_proxy(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
request = Request(self.getURL(server_port, "path/to/resource"))
response = await download_request(download_handler, request)
assert response.status == 200
assert response.url == request.url
assert response.body == b"/path/to/resource"
@deferred_f_from_coro_f
async def test_download_with_proxy_https_timeout(self):
async def test_download_with_proxy_https_timeout(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
if NON_EXISTING_RESOLVABLE:
pytest.skip("Non-existing hosts are resolvable")
http_proxy = self.getURL("")
http_proxy = self.getURL(server_port, "")
domain = "https://no-such-domain.nosuch"
request = Request(domain, meta={"proxy": http_proxy, "download_timeout": 0.2})
with pytest.raises(error.TimeoutError) as exc_info:
await self.download_request(request, Spider("foo"))
await download_request(download_handler, request)
assert domain in exc_info.value.osError
@deferred_f_from_coro_f
async def test_download_with_proxy_without_http_scheme(self):
http_proxy = self.getURL("").replace("http://", "")
async def test_download_with_proxy_without_http_scheme(
self, server_port: int, download_handler: DownloadHandlerProtocol
) -> None:
http_proxy = self.getURL(server_port, "").replace("http://", "")
request = Request("http://example.com", meta={"proxy": http_proxy})
response = await self.download_request(request, Spider("foo"))
response = await download_request(download_handler, request)
assert response.status == 200
assert response.url == request.url
assert response.body == self.expected_http_proxy_request_body

View File

@ -1,12 +1,12 @@
from __future__ import annotations
import asyncio
from contextlib import asynccontextmanager
from gzip import BadGzipFile
from unittest import mock
import pytest
from twisted.internet.defer import Deferred, inlineCallbacks, succeed
from twisted.trial.unittest import TestCase
from twisted.internet.defer import Deferred, succeed
from scrapy.core.downloader.middleware import DownloaderMiddlewareManager
from scrapy.exceptions import _InvalidOutput
@ -17,23 +17,26 @@ from scrapy.utils.python import to_bytes
from scrapy.utils.test import get_crawler, get_from_asyncio_queue
class TestManagerBase(TestCase):
class TestManagerBase:
settings_dict = None
@inlineCallbacks
def setUp(self):
self.crawler = get_crawler(Spider, self.settings_dict)
self.spider = self.crawler._create_spider("foo")
self.mwman = DownloaderMiddlewareManager.from_crawler(self.crawler)
self.crawler.engine = self.crawler._create_engine()
yield self.crawler.engine.open_spider(self.spider)
@inlineCallbacks
def tearDown(self):
yield self.crawler.engine.close_spider(self.spider)
# should be a fixture but async fixtures that use Futures are problematic with pytest-twisted
@asynccontextmanager
async def get_mwman_and_spider(self):
crawler = get_crawler(Spider, self.settings_dict)
spider = crawler._create_spider("foo")
mwman = DownloaderMiddlewareManager.from_crawler(crawler)
crawler.engine = crawler._create_engine()
await crawler.engine.open_spider_async(spider)
yield mwman, spider
await maybe_deferred_to_future(crawler.engine.close_spider(spider))
@staticmethod
async def _download(
self, request: Request, response: Response | None = None
mwman: DownloaderMiddlewareManager,
spider: Spider,
request: Request,
response: Response | None = None,
) -> Response | Request:
"""Executes downloader mw manager's download method and returns
the result (Request or Response) or raises exception in case of
@ -46,7 +49,7 @@ class TestManagerBase(TestCase):
return succeed(response)
return await maybe_deferred_to_future(
self.mwman.download(download_func, request, self.spider)
mwman.download(download_func, request, spider)
)
@ -57,7 +60,8 @@ class TestDefaults(TestManagerBase):
async def test_request_response(self):
req = Request("http://example.com/index.html")
resp = Response(req.url, status=200)
ret = await self._download(req, resp)
async with self.get_mwman_and_spider() as (mwman, spider):
ret = await self._download(mwman, spider, req, resp)
assert isinstance(ret, Response), "Non-response returned"
@deferred_f_from_coro_f
@ -86,7 +90,8 @@ class TestDefaults(TestManagerBase):
"Location": "http://example.com/login",
},
)
ret = await self._download(req, resp)
async with self.get_mwman_and_spider() as (mwman, spider):
ret = await self._download(mwman, spider, req, resp)
assert isinstance(ret, Request), f"Not redirected: {ret!r}"
assert to_bytes(ret.url) == resp.headers["Location"], (
"Not redirected to location header"
@ -108,7 +113,8 @@ class TestDefaults(TestManagerBase):
},
)
with pytest.raises(BadGzipFile):
await self._download(req, resp)
async with self.get_mwman_and_spider() as (mwman, spider):
await self._download(mwman, spider, req, resp)
class TestResponseFromProcessRequest(TestManagerBase):
@ -116,19 +122,19 @@ class TestResponseFromProcessRequest(TestManagerBase):
@deferred_f_from_coro_f
async def test_download_func_not_called(self):
req = Request("http://example.com/index.html")
resp = Response("http://example.com/index.html")
download_func = mock.MagicMock()
class ResponseMiddleware:
def process_request(self, request, spider):
return resp
self.mwman._add_middleware(ResponseMiddleware())
req = Request("http://example.com/index.html")
download_func = mock.MagicMock()
result = await maybe_deferred_to_future(
self.mwman.download(download_func, req, self.spider)
)
async with self.get_mwman_and_spider() as (mwman, spider):
mwman._add_middleware(ResponseMiddleware())
result = await maybe_deferred_to_future(
mwman.download(download_func, req, spider)
)
assert result is resp
assert not download_func.called
@ -138,6 +144,7 @@ class TestResponseFromProcessException(TestManagerBase):
@deferred_f_from_coro_f
async def test_process_response_called(self):
req = Request("http://example.com/index.html")
resp = Response("http://example.com/index.html")
calls = []
@ -153,12 +160,11 @@ class TestResponseFromProcessException(TestManagerBase):
calls.append("process_exception")
return resp
self.mwman._add_middleware(ResponseMiddleware())
req = Request("http://example.com/index.html")
result = await maybe_deferred_to_future(
self.mwman.download(download_func, req, self.spider)
)
async with self.get_mwman_and_spider() as (mwman, spider):
mwman._add_middleware(ResponseMiddleware())
result = await maybe_deferred_to_future(
mwman.download(download_func, req, spider)
)
assert result is resp
assert calls == [
"process_exception",
@ -176,9 +182,10 @@ class TestInvalidOutput(TestManagerBase):
def process_request(self, request, spider):
return 1
self.mwman._add_middleware(InvalidProcessRequestMiddleware())
with pytest.raises(_InvalidOutput):
await self._download(req)
async with self.get_mwman_and_spider() as (mwman, spider):
mwman._add_middleware(InvalidProcessRequestMiddleware())
with pytest.raises(_InvalidOutput):
await self._download(mwman, spider, req)
@deferred_f_from_coro_f
async def test_invalid_process_response(self):
@ -189,9 +196,10 @@ class TestInvalidOutput(TestManagerBase):
def process_response(self, request, response, spider):
return 1
self.mwman._add_middleware(InvalidProcessResponseMiddleware())
with pytest.raises(_InvalidOutput):
await self._download(req)
async with self.get_mwman_and_spider() as (mwman, spider):
mwman._add_middleware(InvalidProcessResponseMiddleware())
with pytest.raises(_InvalidOutput):
await self._download(mwman, spider, req)
@deferred_f_from_coro_f
async def test_invalid_process_exception(self):
@ -205,9 +213,10 @@ class TestInvalidOutput(TestManagerBase):
def process_exception(self, request, exception, spider):
return 1
self.mwman._add_middleware(InvalidProcessExceptionMiddleware())
with pytest.raises(_InvalidOutput):
await self._download(req)
async with self.get_mwman_and_spider() as (mwman, spider):
mwman._add_middleware(InvalidProcessExceptionMiddleware())
with pytest.raises(_InvalidOutput):
await self._download(mwman, spider, req)
class TestMiddlewareUsingDeferreds(TestManagerBase):
@ -215,7 +224,9 @@ class TestMiddlewareUsingDeferreds(TestManagerBase):
@deferred_f_from_coro_f
async def test_deferred(self):
req = Request("http://example.com/index.html")
resp = Response("http://example.com/index.html")
download_func = mock.MagicMock()
class DeferredMiddleware:
def cb(self, result):
@ -227,53 +238,53 @@ class TestMiddlewareUsingDeferreds(TestManagerBase):
d.callback(resp)
return d
self.mwman._add_middleware(DeferredMiddleware())
req = Request("http://example.com/index.html")
download_func = mock.MagicMock()
result = await maybe_deferred_to_future(
self.mwman.download(download_func, req, self.spider)
)
async with self.get_mwman_and_spider() as (mwman, spider):
mwman._add_middleware(DeferredMiddleware())
result = await maybe_deferred_to_future(
mwman.download(download_func, req, spider)
)
assert result is resp
assert not download_func.called
@pytest.mark.usefixtures("reactor_pytest")
class TestMiddlewareUsingCoro(TestManagerBase):
"""Middlewares using asyncio coroutines should work"""
@deferred_f_from_coro_f
async def test_asyncdef(self):
req = Request("http://example.com/index.html")
resp = Response("http://example.com/index.html")
download_func = mock.MagicMock()
class CoroMiddleware:
async def process_request(self, request, spider):
await succeed(42)
return resp
self.mwman._add_middleware(CoroMiddleware())
req = Request("http://example.com/index.html")
download_func = mock.MagicMock()
result = await maybe_deferred_to_future(
self.mwman.download(download_func, req, self.spider)
)
async with self.get_mwman_and_spider() as (mwman, spider):
mwman._add_middleware(CoroMiddleware())
result = await maybe_deferred_to_future(
mwman.download(download_func, req, spider)
)
assert result is resp
assert not download_func.called
@pytest.mark.only_asyncio
@deferred_f_from_coro_f
async def test_asyncdef_asyncio(self):
req = Request("http://example.com/index.html")
resp = Response("http://example.com/index.html")
download_func = mock.MagicMock()
class CoroMiddleware:
async def process_request(self, request, spider):
await asyncio.sleep(0.1)
return await get_from_asyncio_queue(resp)
self.mwman._add_middleware(CoroMiddleware())
req = Request("http://example.com/index.html")
download_func = mock.MagicMock()
result = await maybe_deferred_to_future(
self.mwman.download(download_func, req, self.spider)
)
async with self.get_mwman_and_spider() as (mwman, spider):
mwman._add_middleware(CoroMiddleware())
result = await maybe_deferred_to_future(
mwman.download(download_func, req, spider)
)
assert result is resp
assert not download_func.called

View File

@ -7,7 +7,6 @@ import pytest
from twisted.internet import error
from twisted.internet.defer import Deferred, maybeDeferred
from twisted.python import failure
from twisted.trial import unittest
from scrapy.downloadermiddlewares.robotstxt import RobotsTxtMiddleware
from scrapy.downloadermiddlewares.robotstxt import logger as mw_module_logger
@ -22,13 +21,13 @@ if TYPE_CHECKING:
from scrapy.crawler import Crawler
class TestRobotsTxtMiddleware(unittest.TestCase):
def setUp(self):
class TestRobotsTxtMiddleware:
def setup_method(self):
self.crawler = mock.MagicMock()
self.crawler.settings = Settings()
self.crawler.engine.download = mock.MagicMock()
def tearDown(self):
def teardown_method(self):
del self.crawler
def test_robotstxt_settings(self):
@ -249,8 +248,8 @@ Disallow: /some/randome/page.html
@pytest.mark.skipif(not rerp_available(), reason="Rerp parser is not installed")
class TestRobotsTxtMiddlewareWithRerp(TestRobotsTxtMiddleware):
def setUp(self):
super().setUp()
def setup_method(self):
super().setup_method()
self.crawler.settings.set(
"ROBOTSTXT_PARSER", "scrapy.robotstxt.RerpRobotParser"
)

View File

@ -1,7 +1,6 @@
import time
from twisted.internet.defer import inlineCallbacks
from twisted.trial.unittest import TestCase
from scrapy import Request
from scrapy.core.downloader import Downloader, Slot
@ -49,17 +48,17 @@ class DownloaderSlotsSettingsTestSpider(MetaSpider):
self.times[slot].append(time.time())
class TestCrawl(TestCase):
class TestCrawl:
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
def setUp(self):
def setup_method(self):
self.runner = CrawlerRunner()
@inlineCallbacks

View File

@ -27,7 +27,6 @@ from pydispatch import dispatcher
from testfixtures import LogCapture
from twisted.internet import defer
from twisted.internet.defer import inlineCallbacks
from twisted.trial import unittest
from twisted.web import server, static, util
from scrapy import signals
@ -246,7 +245,7 @@ class CrawlerRun:
self.signals_caught[sig] = signalargs
class TestEngineBase(unittest.TestCase):
class TestEngineBase:
@staticmethod
def _assert_visited_urls(run: CrawlerRun) -> None:
must_be_visited = [

View File

@ -6,7 +6,6 @@ from typing import TYPE_CHECKING
from testfixtures import LogCapture
from twisted.internet.defer import Deferred
from twisted.trial.unittest import TestCase
from scrapy import Request, Spider, signals
from scrapy.utils.defer import deferred_f_from_coro_f, maybe_deferred_to_future
@ -27,7 +26,7 @@ async def sleep(seconds: float = 0.001) -> None:
await maybe_deferred_to_future(deferred)
class TestMain(TestCase):
class TestMain:
@deferred_f_from_coro_f
async def test_sleep(self):
"""Neither asynchronous sleeps on Spider.start() nor the equivalent on
@ -120,16 +119,16 @@ class TestMain(TestCase):
assert actual_urls == expected_urls, f"{actual_urls=} != {expected_urls=}"
class TestRequestSendOrder(TestCase):
class TestRequestSendOrder:
seconds = 0.1 # increase if flaky
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None) # increase if flaky
def request(self, num, response_seconds, download_slots, priority=0):

View File

@ -2,13 +2,12 @@ import pytest
from twisted.conch.telnet import ITelnetProtocol
from twisted.cred import credentials
from twisted.internet.defer import inlineCallbacks
from twisted.trial import unittest
from scrapy.extensions.telnet import TelnetConsole
from scrapy.utils.test import get_crawler
class TestTelnetExtension(unittest.TestCase):
class TestTelnetExtension:
def _get_console_and_portal(self, settings=None):
crawler = get_crawler(settings_dict=settings)
console = TelnetConsole(crawler)

View File

@ -31,7 +31,6 @@ from packaging.version import Version
from testfixtures import LogCapture
from twisted.internet import defer
from twisted.internet.defer import inlineCallbacks
from twisted.trial import unittest
from w3lib.url import file_uri_to_path, path_to_file_uri
from zope.interface import implementer
from zope.interface.verify import verifyObject
@ -164,7 +163,7 @@ class TestFileFeedStorage:
assert storage.path == path
class TestFTPFeedStorage(unittest.TestCase):
class TestFTPFeedStorage:
def get_test_spider(self, settings=None):
class TestSpider(scrapy.Spider):
name = "test_spider"
@ -278,7 +277,7 @@ class TestBlockingFeedStorage:
@pytest.mark.requires_boto3
class TestS3FeedStorage(unittest.TestCase):
class TestS3FeedStorage:
def test_parse_credentials(self):
aws_credentials = {
"AWS_ACCESS_KEY_ID": "settings_key",
@ -507,7 +506,7 @@ class TestS3FeedStorage(unittest.TestCase):
assert "S3 does not support appending to files" in str(log)
class TestGCSFeedStorage(unittest.TestCase):
class TestGCSFeedStorage:
def test_parse_settings(self):
try:
from google.cloud.storage import Client # noqa: F401,PLC0415
@ -661,7 +660,7 @@ class LogOnStoreFileStorage:
file.close()
class TestFeedExportBase(ABC, unittest.TestCase):
class TestFeedExportBase(ABC):
mockserver: MockServer
class MyItem(scrapy.Item):
@ -679,18 +678,18 @@ class TestFeedExportBase(ABC, unittest.TestCase):
return Path(self.temp_dir, inter_dir, filename)
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
def setUp(self):
def setup_method(self):
self.temp_dir = tempfile.mkdtemp()
def tearDown(self):
def teardown_method(self):
shutil.rmtree(self.temp_dir, ignore_errors=True)
async def exported_data(
@ -735,7 +734,7 @@ class TestFeedExportBase(ABC, unittest.TestCase):
await self.assertExportedMarshal(items, rows, settings)
await self.assertExportedMultiple(items, rows, settings)
async def assertExportedCsv(
async def assertExportedCsv( # noqa: B027
self,
items: Iterable[Any],
header: Iterable[str],
@ -744,7 +743,7 @@ class TestFeedExportBase(ABC, unittest.TestCase):
) -> None:
pass
async def assertExportedJsonLines(
async def assertExportedJsonLines( # noqa: B027
self,
items: Iterable[Any],
rows: Iterable[dict[str, Any]],
@ -752,7 +751,7 @@ class TestFeedExportBase(ABC, unittest.TestCase):
) -> None:
pass
async def assertExportedXml(
async def assertExportedXml( # noqa: B027
self,
items: Iterable[Any],
rows: Iterable[dict[str, Any]],
@ -760,7 +759,7 @@ class TestFeedExportBase(ABC, unittest.TestCase):
) -> None:
pass
async def assertExportedMultiple(
async def assertExportedMultiple( # noqa: B027
self,
items: Iterable[Any],
rows: Iterable[dict[str, Any]],
@ -768,7 +767,7 @@ class TestFeedExportBase(ABC, unittest.TestCase):
) -> None:
pass
async def assertExportedPickle(
async def assertExportedPickle( # noqa: B027
self,
items: Iterable[Any],
rows: Iterable[dict[str, Any]],
@ -776,7 +775,7 @@ class TestFeedExportBase(ABC, unittest.TestCase):
) -> None:
pass
async def assertExportedMarshal(
async def assertExportedMarshal( # noqa: B027
self,
items: Iterable[Any],
rows: Iterable[dict[str, Any]],

View File

@ -3,16 +3,15 @@ from __future__ import annotations
import json
import random
import re
import shutil
import string
from ipaddress import IPv4Address
from pathlib import Path
from tempfile import mkdtemp
from typing import TYPE_CHECKING, Any, Callable
from typing import TYPE_CHECKING, Any, Callable, cast
from unittest import mock
from urllib.parse import urlencode
import pytest
from pytest_twisted import async_yield_fixture
from twisted.internet.defer import (
CancelledError,
Deferred,
@ -22,7 +21,6 @@ from twisted.internet.defer import (
from twisted.internet.endpoints import SSL4ClientEndpoint, SSL4ServerEndpoint
from twisted.internet.error import TimeoutError as TxTimeoutError
from twisted.internet.ssl import Certificate, PrivateCertificate, optionsForClientTLS
from twisted.trial.unittest import TestCase
from twisted.web.client import URI, ResponseFailed
from twisted.web.http import H2_ENABLED
from twisted.web.http import Request as TxRequest
@ -40,7 +38,9 @@ from scrapy.utils.defer import (
from tests.mockserver import LeafResource, Status, ssl_context_factory
if TYPE_CHECKING:
from collections.abc import Coroutine
from collections.abc import AsyncGenerator, Coroutine, Generator
from scrapy.core.http2.protocol import H2ClientProtocol
def generate_random_string(size: int) -> str:
@ -178,24 +178,24 @@ class RequestHeaders(LeafResource):
return bytes(json.dumps(headers), "utf-8")
def get_client_certificate(
key_file: Path, certificate_file: Path
) -> PrivateCertificate:
pem = key_file.read_text(encoding="utf-8") + certificate_file.read_text(
encoding="utf-8"
)
return PrivateCertificate.loadPEM(pem)
def make_request_dfd(client: H2ClientProtocol, request: Request) -> Deferred[Response]:
return client.request(request, DummySpider())
async def make_request(client: H2ClientProtocol, request: Request) -> Response:
return await maybe_deferred_to_future(make_request_dfd(client, request))
@pytest.mark.skipif(not H2_ENABLED, reason="HTTP/2 support in Twisted is not enabled")
class TestHttps2ClientProtocol(TestCase):
class TestHttps2ClientProtocol:
scheme = "https"
host = "localhost"
key_file = Path(__file__).parent / "keys" / "localhost.key"
certificate_file = Path(__file__).parent / "keys" / "localhost.crt"
def _init_resource(self):
self.temp_directory = mkdtemp()
r = File(self.temp_directory)
@pytest.fixture
def site(self, tmp_path):
r = File(str(tmp_path))
r.putChild(b"get-data-html-small", GetDataHtmlSmall())
r.putChild(b"get-data-html-large", GetDataHtmlLarge())
@ -208,72 +208,65 @@ class TestHttps2ClientProtocol(TestCase):
r.putChild(b"query-params", QueryParams())
r.putChild(b"timeout", TimeoutResponse())
r.putChild(b"request-headers", RequestHeaders())
return r
return Site(r, timeout=None)
@inlineCallbacks
def setUp(self):
@async_yield_fixture
async def server_port(self, site: Site) -> AsyncGenerator[int]:
from twisted.internet import reactor
# Initialize resource tree
root = self._init_resource()
self.site = Site(root, timeout=None)
# Start server for testing
self.hostname = "localhost"
context_factory = ssl_context_factory(
str(self.key_file), str(self.certificate_file)
)
server_endpoint = SSL4ServerEndpoint(
reactor, 0, context_factory, interface=self.hostname
reactor, 0, context_factory, interface=self.host
)
self.server = yield server_endpoint.listen(self.site)
self.port_number = self.server.getHost().port
server = await server_endpoint.listen(site)
# Connect H2 client with server
self.client_certificate = get_client_certificate(
self.key_file, self.certificate_file
)
client_options = optionsForClientTLS(
hostname=self.hostname,
trustRoot=self.client_certificate,
acceptableProtocols=[b"h2"],
)
uri = URI.fromBytes(bytes(self.get_url("/"), "utf-8"))
yield server.getHost().port
self.conn_closed_deferred = Deferred()
await server.stopListening()
@pytest.fixture
def client_certificate(self) -> PrivateCertificate:
pem = self.key_file.read_text(
encoding="utf-8"
) + self.certificate_file.read_text(encoding="utf-8")
return PrivateCertificate.loadPEM(pem)
@async_yield_fixture
async def client(
self, server_port: int, client_certificate: PrivateCertificate
) -> AsyncGenerator[H2ClientProtocol]:
from twisted.internet import reactor
from scrapy.core.http2.protocol import H2ClientFactory # noqa: PLC0415
h2_client_factory = H2ClientFactory(uri, Settings(), self.conn_closed_deferred)
client_endpoint = SSL4ClientEndpoint(
reactor, self.hostname, self.port_number, client_options
client_options = optionsForClientTLS(
hostname=self.host,
trustRoot=client_certificate,
acceptableProtocols=[b"h2"],
)
self.client = yield client_endpoint.connect(h2_client_factory)
uri = URI.fromBytes(bytes(self.get_url(server_port, "/"), "utf-8"))
h2_client_factory = H2ClientFactory(uri, Settings(), Deferred())
client_endpoint = SSL4ClientEndpoint(
reactor, self.host, server_port, client_options
)
client = await client_endpoint.connect(h2_client_factory)
@inlineCallbacks
def tearDown(self):
if self.client.connected:
yield self.client.transport.loseConnection()
yield self.client.transport.abortConnection()
yield self.server.stopListening()
shutil.rmtree(self.temp_directory)
self.conn_closed_deferred = None
yield client
def get_url(self, path: str) -> str:
if client.connected:
client.transport.loseConnection()
client.transport.abortConnection()
def get_url(self, portno: int, path: str) -> str:
"""
:param path: Should have / at the starting compulsorily if not empty
:return: Complete url
"""
assert len(path) > 0
assert path[0] == "/" or path[0] == "&"
return f"{self.scheme}://{self.hostname}:{self.port_number}{path}"
async def make_request(self, request: Request) -> Response:
return await maybe_deferred_to_future(self.make_request_dfd(request))
def make_request_dfd(self, request: Request) -> Deferred[Response]:
return self.client.request(request, DummySpider())
return f"{self.scheme}://{self.host}:{portno}{path}"
@staticmethod
async def _check_repeat(
@ -287,9 +280,13 @@ class TestHttps2ClientProtocol(TestCase):
await maybe_deferred_to_future(DeferredList(d_list, fireOnOneErrback=True))
async def _check_GET(
self, request: Request, expected_body: bytes, expected_status: int
self,
client: H2ClientProtocol,
request: Request,
expected_body: bytes,
expected_status: int,
) -> None:
response = await self.make_request(request)
response = await make_request(client, request)
assert response.status == expected_status
assert response.body == expected_body
assert response.request == request
@ -300,43 +297,62 @@ class TestHttps2ClientProtocol(TestCase):
assert len(response.body) == content_length
@deferred_f_from_coro_f
async def test_GET_small_body(self):
request = Request(self.get_url("/get-data-html-small"))
await self._check_GET(request, Data.HTML_SMALL, 200)
async def test_GET_small_body(
self, server_port: int, client: H2ClientProtocol
) -> None:
request = Request(self.get_url(server_port, "/get-data-html-small"))
await self._check_GET(client, request, Data.HTML_SMALL, 200)
@deferred_f_from_coro_f
async def test_GET_large_body(self):
request = Request(self.get_url("/get-data-html-large"))
await self._check_GET(request, Data.HTML_LARGE, 200)
async def test_GET_large_body(
self, server_port: int, client: H2ClientProtocol
) -> None:
request = Request(self.get_url(server_port, "/get-data-html-large"))
await self._check_GET(client, request, Data.HTML_LARGE, 200)
async def _check_GET_x10(
self, request: Request, expected_body: bytes, expected_status: int
self,
client: H2ClientProtocol,
request: Request,
expected_body: bytes,
expected_status: int,
) -> None:
async def get_coro() -> None:
await self._check_GET(request, expected_body, expected_status)
await self._check_GET(client, request, expected_body, expected_status)
await self._check_repeat(get_coro, 10)
@deferred_f_from_coro_f
async def test_GET_small_body_x10(self):
async def test_GET_small_body_x10(
self, server_port: int, client: H2ClientProtocol
) -> None:
await self._check_GET_x10(
Request(self.get_url("/get-data-html-small")), Data.HTML_SMALL, 200
client,
Request(self.get_url(server_port, "/get-data-html-small")),
Data.HTML_SMALL,
200,
)
@deferred_f_from_coro_f
async def test_GET_large_body_x10(self):
async def test_GET_large_body_x10(
self, server_port: int, client: H2ClientProtocol
) -> None:
await self._check_GET_x10(
Request(self.get_url("/get-data-html-large")), Data.HTML_LARGE, 200
client,
Request(self.get_url(server_port, "/get-data-html-large")),
Data.HTML_LARGE,
200,
)
@staticmethod
async def _check_POST_json(
self,
client: H2ClientProtocol,
request: Request,
expected_request_body: dict[str, str],
expected_extra_data: str,
expected_status: int,
) -> None:
response = await self.make_request(request)
response = await make_request(client, request)
assert response.status == expected_status
assert response.request == request
@ -369,22 +385,30 @@ class TestHttps2ClientProtocol(TestCase):
assert request_headers[k_str] == str(v[0], "utf-8")
@deferred_f_from_coro_f
async def test_POST_small_json(self):
async def test_POST_small_json(
self, server_port: int, client: H2ClientProtocol
) -> None:
request = JsonRequest(
url=self.get_url("/post-data-json-small"),
url=self.get_url(server_port, "/post-data-json-small"),
method="POST",
data=Data.JSON_SMALL,
)
await self._check_POST_json(request, Data.JSON_SMALL, Data.EXTRA_SMALL, 200)
await self._check_POST_json(
client, request, Data.JSON_SMALL, Data.EXTRA_SMALL, 200
)
@deferred_f_from_coro_f
async def test_POST_large_json(self):
async def test_POST_large_json(
self, server_port: int, client: H2ClientProtocol
) -> None:
request = JsonRequest(
url=self.get_url("/post-data-json-large"),
url=self.get_url(server_port, "/post-data-json-large"),
method="POST",
data=Data.JSON_LARGE,
)
await self._check_POST_json(request, Data.JSON_LARGE, Data.EXTRA_LARGE, 200)
await self._check_POST_json(
client, request, Data.JSON_LARGE, Data.EXTRA_LARGE, 200
)
async def _check_POST_json_x10(self, *args, **kwargs):
async def get_coro() -> None:
@ -393,48 +417,63 @@ class TestHttps2ClientProtocol(TestCase):
await self._check_repeat(get_coro, 10)
@deferred_f_from_coro_f
async def test_POST_small_json_x10(self):
async def test_POST_small_json_x10(
self, server_port: int, client: H2ClientProtocol
) -> None:
request = JsonRequest(
url=self.get_url("/post-data-json-small"),
url=self.get_url(server_port, "/post-data-json-small"),
method="POST",
data=Data.JSON_SMALL,
)
await self._check_POST_json_x10(request, Data.JSON_SMALL, Data.EXTRA_SMALL, 200)
await self._check_POST_json_x10(
client, request, Data.JSON_SMALL, Data.EXTRA_SMALL, 200
)
@deferred_f_from_coro_f
async def test_POST_large_json_x10(self):
async def test_POST_large_json_x10(
self, server_port: int, client: H2ClientProtocol
) -> None:
request = JsonRequest(
url=self.get_url("/post-data-json-large"),
url=self.get_url(server_port, "/post-data-json-large"),
method="POST",
data=Data.JSON_LARGE,
)
await self._check_POST_json_x10(request, Data.JSON_LARGE, Data.EXTRA_LARGE, 200)
await self._check_POST_json_x10(
client, request, Data.JSON_LARGE, Data.EXTRA_LARGE, 200
)
@inlineCallbacks
def test_invalid_negotiated_protocol(self):
def test_invalid_negotiated_protocol(
self, server_port: int, client: H2ClientProtocol
) -> Generator[Deferred[Any], Any, None]:
with mock.patch(
"scrapy.core.http2.protocol.PROTOCOL_NAME", return_value=b"not-h2"
):
request = Request(url=self.get_url("/status?n=200"))
request = Request(url=self.get_url(server_port, "/status?n=200"))
with pytest.raises(ResponseFailed):
yield self.make_request_dfd(request)
yield make_request_dfd(client, request)
@inlineCallbacks
def test_cancel_request(self):
request = Request(url=self.get_url("/get-data-html-large"))
d = self.make_request_dfd(request)
def test_cancel_request(
self, server_port: int, client: H2ClientProtocol
) -> Generator[Deferred[Any], Any, None]:
request = Request(url=self.get_url(server_port, "/get-data-html-large"))
d = make_request_dfd(client, request)
d.cancel()
response = yield d
response = cast("Response", (yield d))
assert response.status == 499
assert response.request == request
@deferred_f_from_coro_f
async def test_download_maxsize_exceeded(self):
async def test_download_maxsize_exceeded(
self, server_port: int, client: H2ClientProtocol
) -> None:
request = Request(
url=self.get_url("/get-data-html-large"), meta={"download_maxsize": 1000}
url=self.get_url(server_port, "/get-data-html-large"),
meta={"download_maxsize": 1000},
)
with pytest.raises(CancelledError) as exc_info:
await self.make_request(request)
await make_request(client, request)
error_pattern = re.compile(
rf"Cancelling download of {request.url}: received response "
rf"size \(\d*\) larger than download max size \(1000\)"
@ -442,14 +481,16 @@ class TestHttps2ClientProtocol(TestCase):
assert len(re.findall(error_pattern, str(exc_info.value))) == 1
@inlineCallbacks
def test_received_dataloss_response(self):
def test_received_dataloss_response(
self, server_port: int, client: H2ClientProtocol
) -> Generator[Deferred[Any], Any, None]:
"""In case when value of Header Content-Length != len(Received Data)
ProtocolError is raised"""
from h2.exceptions import InvalidBodyLengthError # noqa: PLC0415
request = Request(url=self.get_url("/dataloss"))
request = Request(url=self.get_url(server_port, "/dataloss"))
with pytest.raises(ResponseFailed) as exc_info:
yield self.make_request_dfd(request)
yield make_request_dfd(client, request)
assert len(exc_info.value.reasons) > 0
assert any(
isinstance(error, InvalidBodyLengthError)
@ -457,42 +498,62 @@ class TestHttps2ClientProtocol(TestCase):
)
@deferred_f_from_coro_f
async def test_missing_content_length_header(self):
request = Request(url=self.get_url("/no-content-length-header"))
response = await self.make_request(request)
async def test_missing_content_length_header(
self, server_port: int, client: H2ClientProtocol
) -> None:
request = Request(url=self.get_url(server_port, "/no-content-length-header"))
response = await make_request(client, request)
assert response.status == 200
assert response.body == Data.NO_CONTENT_LENGTH
assert response.request == request
assert "Content-Length" not in response.headers
async def _check_log_warnsize(
self, request: Request, warn_pattern: re.Pattern[str], expected_body: bytes
self,
client: H2ClientProtocol,
request: Request,
warn_pattern: re.Pattern[str],
expected_body: bytes,
caplog: pytest.LogCaptureFixture,
) -> None:
with self.assertLogs("scrapy.core.http2.stream", level="WARNING") as cm:
response = await self.make_request(request)
assert response.status == 200
assert response.request == request
assert response.body == expected_body
with caplog.at_level("WARNING", "scrapy.core.http2.stream"):
response = await make_request(client, request)
assert response.status == 200
assert response.request == request
assert response.body == expected_body
# Check the warning is raised only once for this request
assert sum(len(re.findall(warn_pattern, log)) for log in cm.output) == 1
# Check the warning is raised only once for this request
assert len(re.findall(warn_pattern, caplog.text)) == 1
@deferred_f_from_coro_f
async def test_log_expected_warnsize(self):
async def test_log_expected_warnsize(
self,
server_port: int,
client: H2ClientProtocol,
caplog: pytest.LogCaptureFixture,
) -> None:
request = Request(
url=self.get_url("/get-data-html-large"), meta={"download_warnsize": 1000}
url=self.get_url(server_port, "/get-data-html-large"),
meta={"download_warnsize": 1000},
)
warn_pattern = re.compile(
rf"Expected response size \(\d*\) larger than "
rf"download warn size \(1000\) in request {request}"
)
await self._check_log_warnsize(request, warn_pattern, Data.HTML_LARGE)
await self._check_log_warnsize(
client, request, warn_pattern, Data.HTML_LARGE, caplog
)
@deferred_f_from_coro_f
async def test_log_received_warnsize(self):
async def test_log_received_warnsize(
self,
server_port: int,
client: H2ClientProtocol,
caplog: pytest.LogCaptureFixture,
) -> None:
request = Request(
url=self.get_url("/no-content-length-header"),
url=self.get_url(server_port, "/no-content-length-header"),
meta={"download_warnsize": 10},
)
warn_pattern = re.compile(
@ -500,23 +561,32 @@ class TestHttps2ClientProtocol(TestCase):
rf"warn size \(10\) in request {request}"
)
await self._check_log_warnsize(request, warn_pattern, Data.NO_CONTENT_LENGTH)
await self._check_log_warnsize(
client, request, warn_pattern, Data.NO_CONTENT_LENGTH, caplog
)
@deferred_f_from_coro_f
async def test_max_concurrent_streams(self):
async def test_max_concurrent_streams(
self, server_port: int, client: H2ClientProtocol
) -> None:
"""Send 500 requests at one to check if we can handle
very large number of request.
"""
async def get_coro() -> None:
await self._check_GET(
Request(self.get_url("/get-data-html-small")), Data.HTML_SMALL, 200
client,
Request(self.get_url(server_port, "/get-data-html-small")),
Data.HTML_SMALL,
200,
)
await self._check_repeat(get_coro, 500)
@inlineCallbacks
def test_inactive_stream(self):
def test_inactive_stream(
self, server_port: int, client: H2ClientProtocol
) -> Generator[Deferred[Any], Any, None]:
"""Here we send 110 requests considering the MAX_CONCURRENT_STREAMS
by default is 100. After sending the first 100 requests we close the
connection."""
@ -533,38 +603,47 @@ class TestHttps2ClientProtocol(TestCase):
# Send 100 request (we do not check the result)
for _ in range(100):
d = self.make_request_dfd(Request(self.get_url("/get-data-html-small")))
d = make_request_dfd(
client, Request(self.get_url(server_port, "/get-data-html-small"))
)
d.addBoth(lambda _: None)
d_list.append(d)
# Now send 10 extra request and save the response deferred in a list
for _ in range(10):
d = self.make_request_dfd(Request(self.get_url("/get-data-html-small")))
d = make_request_dfd(
client, Request(self.get_url(server_port, "/get-data-html-small"))
)
d.addCallback(lambda _: pytest.fail("This request should have failed"))
d.addErrback(assert_inactive_stream)
d_list.append(d)
# Close the connection now to fire all the extra 10 requests errback
# with InactiveStreamClosed
self.client.transport.loseConnection()
assert client.transport
client.transport.loseConnection()
yield DeferredList(d_list, consumeErrors=True, fireOnOneErrback=True)
@deferred_f_from_coro_f
async def test_invalid_request_type(self):
async def test_invalid_request_type(self, client: H2ClientProtocol):
with pytest.raises(TypeError):
await self.make_request("https://InvalidDataTypePassed.com")
await make_request(client, "https://InvalidDataTypePassed.com") # type: ignore[arg-type]
@deferred_f_from_coro_f
async def test_query_parameters(self):
async def test_query_parameters(
self, server_port: int, client: H2ClientProtocol
) -> None:
params = {
"a": generate_random_string(20),
"b": generate_random_string(20),
"c": generate_random_string(20),
"d": generate_random_string(20),
}
request = Request(self.get_url(f"/query-params?{urlencode(params)}"))
response = await self.make_request(request)
request = Request(
self.get_url(server_port, f"/query-params?{urlencode(params)}")
)
response = await make_request(client, request)
content_encoding_header = response.headers[b"Content-Encoding"]
assert content_encoding_header is not None
content_encoding = str(content_encoding_header, "utf-8")
@ -572,62 +651,78 @@ class TestHttps2ClientProtocol(TestCase):
assert data == params
@deferred_f_from_coro_f
async def test_status_codes(self):
async def test_status_codes(
self, server_port: int, client: H2ClientProtocol
) -> None:
for status in [200, 404]:
request = Request(self.get_url(f"/status?n={status}"))
response = await self.make_request(request)
request = Request(self.get_url(server_port, f"/status?n={status}"))
response = await make_request(client, request)
assert response.status == status
@deferred_f_from_coro_f
async def test_response_has_correct_certificate_ip_address(self):
request = Request(self.get_url("/status?n=200"))
response = await self.make_request(request)
async def test_response_has_correct_certificate_ip_address(
self,
server_port: int,
client: H2ClientProtocol,
client_certificate: PrivateCertificate,
) -> None:
request = Request(self.get_url(server_port, "/status?n=200"))
response = await make_request(client, request)
assert response.request == request
assert isinstance(response.certificate, Certificate)
assert response.certificate.original is not None
assert response.certificate.getIssuer() == self.client_certificate.getIssuer()
assert response.certificate.getIssuer() == client_certificate.getIssuer()
assert response.certificate.getPublicKey().matches(
self.client_certificate.getPublicKey()
client_certificate.getPublicKey()
)
assert isinstance(response.ip_address, IPv4Address)
assert str(response.ip_address) == "127.0.0.1"
async def _check_invalid_netloc(self, url: str) -> None:
@staticmethod
async def _check_invalid_netloc(client: H2ClientProtocol, url: str) -> None:
from scrapy.core.http2.stream import InvalidHostname # noqa: PLC0415
request = Request(url)
with pytest.raises(InvalidHostname) as exc_info:
await self.make_request(request)
await make_request(client, request)
error_msg = str(exc_info.value)
assert "localhost" in error_msg
assert "127.0.0.1" in error_msg
assert str(request) in error_msg
@deferred_f_from_coro_f
async def test_invalid_hostname(self):
await self._check_invalid_netloc("https://notlocalhost.notlocalhostdomain")
async def test_invalid_hostname(self, client: H2ClientProtocol) -> None:
await self._check_invalid_netloc(
client, "https://notlocalhost.notlocalhostdomain"
)
@deferred_f_from_coro_f
async def test_invalid_host_port(self):
port = self.port_number + 1
await self._check_invalid_netloc(f"https://127.0.0.1:{port}")
async def test_invalid_host_port(
self, server_port: int, client: H2ClientProtocol
) -> None:
port = server_port + 1
await self._check_invalid_netloc(client, f"https://127.0.0.1:{port}")
@deferred_f_from_coro_f
async def test_connection_stays_with_invalid_requests(self):
await maybe_deferred_to_future(self.test_invalid_hostname())
await maybe_deferred_to_future(self.test_invalid_host_port())
await maybe_deferred_to_future(self.test_GET_small_body())
await maybe_deferred_to_future(self.test_POST_small_json())
async def test_connection_stays_with_invalid_requests(
self, server_port: int, client: H2ClientProtocol
):
await maybe_deferred_to_future(self.test_invalid_hostname(client))
await maybe_deferred_to_future(self.test_invalid_host_port(server_port, client))
await maybe_deferred_to_future(self.test_GET_small_body(server_port, client))
await maybe_deferred_to_future(self.test_POST_small_json(server_port, client))
@inlineCallbacks
def test_connection_timeout(self):
request = Request(self.get_url("/timeout"))
def test_connection_timeout(
self, server_port: int, client: H2ClientProtocol
) -> Generator[Deferred[Any], Any, None]:
request = Request(self.get_url(server_port, "/timeout"))
# Update the timer to 1s to test connection timeout
self.client.setTimeout(1)
client.setTimeout(1)
with pytest.raises(ResponseFailed) as exc_info:
yield self.make_request_dfd(request)
yield make_request_dfd(client, request)
for err in exc_info.value.reasons:
from scrapy.core.http2.protocol import H2ClientProtocol # noqa: PLC0415
@ -642,18 +737,20 @@ class TestHttps2ClientProtocol(TestCase):
pytest.fail("No TimeoutError raised.")
@deferred_f_from_coro_f
async def test_request_headers_received(self):
async def test_request_headers_received(
self, server_port: int, client: H2ClientProtocol
) -> None:
request = Request(
self.get_url("/request-headers"),
self.get_url(server_port, "/request-headers"),
headers={"header-1": "header value 1", "header-2": "header value 2"},
)
response = await self.make_request(request)
response = await make_request(client, request)
assert response.status == 200
assert response.request == request
response_headers = json.loads(str(response.body, "utf-8"))
assert isinstance(response_headers, dict)
for k, v in request.headers.items():
k, v = str(k, "utf-8"), str(v[0], "utf-8")
assert k in response_headers
assert v == response_headers[k]
k_decoded, v_decoded = str(k, "utf-8"), str(v[0], "utf-8")
assert k_decoded in response_headers
assert v_decoded == response_headers[k_decoded]

View File

@ -1466,12 +1466,6 @@ class TestJsonRequest(TestRequest):
b"Accept": [b"application/json, text/javascript, */*; q=0.01"],
}
def setup_method(self):
warnings.simplefilter("always")
def teardown_method(self):
warnings.resetwarnings()
def test_data(self):
r1 = self.request_class(url="http://www.example.com/")
assert r1.body == b""

View File

@ -4,7 +4,6 @@ import pytest
from testfixtures import LogCapture
from twisted.internet.defer import inlineCallbacks
from twisted.python.failure import Failure
from twisted.trial.unittest import TestCase
from scrapy.exceptions import DropItem
from scrapy.http import Request, Response
@ -254,17 +253,17 @@ class DropSomeItemsPipeline:
self.drop = True
class TestShowOrSkipMessages(TestCase):
class TestShowOrSkipMessages:
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
def setUp(self):
def setup_method(self):
self.base_settings = {
"LOG_LEVEL": "DEBUG",
"ITEM_PIPELINES": {

View File

@ -8,7 +8,6 @@ from typing import TYPE_CHECKING, Any
import pytest
from testfixtures import LogCapture
from twisted.internet.defer import inlineCallbacks
from twisted.trial.unittest import TestCase
from w3lib.url import add_or_replace_parameter
from scrapy import Spider, signals
@ -58,7 +57,7 @@ class RedirectedMediaDownloadSpider(MediaDownloadSpider):
)
class TestFileDownloadCrawl(TestCase):
class TestFileDownloadCrawl:
pipeline_class = "scrapy.pipelines.files.FilesPipeline"
store_setting_key = "FILES_STORE"
media_key = "files"
@ -70,15 +69,15 @@ class TestFileDownloadCrawl(TestCase):
}
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
def setUp(self):
def setup_method(self):
# prepare a directory for storing files
self.tmpmediastore = Path(mkdtemp())
self.settings = {
@ -87,7 +86,7 @@ class TestFileDownloadCrawl(TestCase):
}
self.items = []
def tearDown(self):
def teardown_method(self):
shutil.rmtree(self.tmpmediastore)
self.items = []

View File

@ -19,7 +19,6 @@ import attr
import pytest
from itemadapter import ItemAdapter
from twisted.internet.defer import inlineCallbacks
from twisted.trial import unittest
from scrapy.http import Request, Response
from scrapy.item import Field, Item
@ -75,8 +74,8 @@ def get_ftp_content_and_delete(
return b"".join(ftp_data)
class TestFilesPipeline(unittest.TestCase):
def setUp(self):
class TestFilesPipeline:
def setup_method(self):
self.tempdir = mkdtemp()
settings_dict = {"FILES_STORE": self.tempdir}
crawler = get_crawler(spidercls=None, settings_dict=settings_dict)
@ -84,7 +83,7 @@ class TestFilesPipeline(unittest.TestCase):
self.pipeline.download_func = _mocked_download_func
self.pipeline.open_spider(None)
def tearDown(self):
def teardown_method(self):
rmtree(self.tempdir)
def test_file_path(self):
@ -538,7 +537,7 @@ class TestFilesPipelineCustomSettings:
@pytest.mark.requires_botocore
class TestS3FilesStore(unittest.TestCase):
class TestS3FilesStore:
@inlineCallbacks
def test_persist(self):
bucket = "mybucket"
@ -615,7 +614,7 @@ class TestS3FilesStore(unittest.TestCase):
@pytest.mark.skipif(
"GCS_PROJECT_ID" not in os.environ, reason="GCS_PROJECT_ID not found"
)
class TestGCSFilesStore(unittest.TestCase):
class TestGCSFilesStore:
@inlineCallbacks
def test_persist(self):
uri = os.environ.get("GCS_TEST_FILE_URI")
@ -667,7 +666,7 @@ class TestGCSFilesStore(unittest.TestCase):
store.bucket.get_blob.assert_called_with(expected_blob_path)
class TestFTPFileStore(unittest.TestCase):
class TestFTPFileStore:
@inlineCallbacks
def test_persist(self):
data = b"TestFTPFilesStore: \xe2\x98\x83"

View File

@ -6,7 +6,6 @@ import pytest
from testfixtures import LogCapture
from twisted.internet.defer import Deferred, inlineCallbacks
from twisted.python.failure import Failure
from twisted.trial import unittest
from scrapy import signals
from scrapy.exceptions import ScrapyDeprecationWarning
@ -43,11 +42,11 @@ class UserDefinedPipeline(MediaPipeline):
return ""
class TestBaseMediaPipeline(unittest.TestCase):
class TestBaseMediaPipeline:
pipeline_class = UserDefinedPipeline
settings = None
def setUp(self):
def setup_method(self):
spider_cls = Spider
self.spider = spider_cls("media.com")
crawler = get_crawler(spider_cls, self.settings)
@ -57,7 +56,7 @@ class TestBaseMediaPipeline(unittest.TestCase):
self.info = self.pipe.spiderinfo
self.fingerprint = crawler.request_fingerprinter.fingerprint
def tearDown(self):
def teardown_method(self):
for name, signal in vars(signals).items():
if not name.startswith("_"):
disconnect_all(signal)
@ -550,13 +549,13 @@ class MediaFailedFailurePipeline(MockedMediaPipeline):
return failure # deprecated
class TestMediaFailedFailure(unittest.TestCase):
class TestMediaFailedFailure:
"""Test that media_failed() can return a failure instead of raising."""
pipeline_class = MediaFailedFailurePipeline
settings = None
def setUp(self):
def setup_method(self):
spider_cls = Spider
self.spider = spider_cls("media.com")
crawler = get_crawler(spider_cls, self.settings)
@ -566,7 +565,7 @@ class TestMediaFailedFailure(unittest.TestCase):
self.info = self.pipe.spiderinfo
self.fingerprint = crawler.request_fingerprinter.fingerprint
def tearDown(self):
def teardown_method(self):
for name, signal in vars(signals).items():
if not name.startswith("_"):
disconnect_all(signal)

View File

@ -2,7 +2,6 @@ import asyncio
import pytest
from twisted.internet.defer import Deferred, inlineCallbacks
from twisted.trial import unittest
from scrapy import Request, Spider, signals
from scrapy.utils.defer import deferred_to_future, maybe_deferred_to_future
@ -75,14 +74,14 @@ class ItemSpider(Spider):
return {"field": 42}
class TestPipeline(unittest.TestCase):
class TestPipeline:
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
def _on_item_scraped(self, item):

View File

@ -9,7 +9,6 @@ from urllib.parse import urlsplit, urlunsplit
import pytest
from testfixtures import LogCapture
from twisted.internet.defer import inlineCallbacks
from twisted.trial.unittest import TestCase
from scrapy.http import Request
from scrapy.utils.test import get_crawler
@ -62,17 +61,17 @@ def _wrong_credentials(proxy_url):
return urlunsplit(bad_auth_proxy)
class TestProxyConnect(TestCase):
class TestProxyConnect:
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
def setUp(self):
def setup_method(self):
try:
import mitmproxy # noqa: F401,PLC0415
except ImportError:
@ -85,7 +84,7 @@ class TestProxyConnect(TestCase):
os.environ["https_proxy"] = proxy_url
os.environ["http_proxy"] = proxy_url
def tearDown(self):
def teardown_method(self):
self._proxy.stop()
os.environ = self._oldenv

View File

@ -1,6 +1,5 @@
from testfixtures import LogCapture
from twisted.internet.defer import inlineCallbacks
from twisted.trial.unittest import TestCase
from scrapy import Request, signals
from scrapy.http.response import Response
@ -56,14 +55,14 @@ class AlternativeCallbacksMiddleware:
return response.replace(request=new_request)
class TestCrawl(TestCase):
class TestCrawl:
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
@inlineCallbacks

View File

@ -1,6 +1,5 @@
from testfixtures import LogCapture
from twisted.internet.defer import inlineCallbacks
from twisted.trial.unittest import TestCase
from scrapy.http import Request
from scrapy.utils.test import get_crawler
@ -149,14 +148,14 @@ class KeywordArgumentsSpider(MockServerSpider):
self.crawler.stats.inc_value("boolean_checks", 1)
class TestCallbackKeywordArguments(TestCase):
class TestCallbackKeywordArguments:
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
@inlineCallbacks

View File

@ -1,5 +1,4 @@
from twisted.internet.defer import inlineCallbacks
from twisted.trial.unittest import TestCase
from scrapy.signals import request_left_downloader
from scrapy.spiders import Spider
@ -24,14 +23,14 @@ class SignalCatcherSpider(Spider):
self.caught_times += 1
class TestCatching(TestCase):
class TestCatching:
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
@inlineCallbacks

View File

@ -9,7 +9,6 @@ from typing import Any, NamedTuple
import pytest
from twisted.internet.defer import inlineCallbacks
from twisted.trial.unittest import TestCase
from scrapy.core.downloader import Downloader
from scrapy.core.scheduler import BaseScheduler, Scheduler
@ -353,8 +352,8 @@ class StartUrlsSpider(Spider):
pass
class TestIntegrationWithDownloaderAwareInMemory(TestCase):
def setUp(self):
class TestIntegrationWithDownloaderAwareInMemory:
def setup_method(self):
self.crawler = get_crawler(
spidercls=StartUrlsSpider,
settings_dict={
@ -363,10 +362,6 @@ class TestIntegrationWithDownloaderAwareInMemory(TestCase):
},
)
@inlineCallbacks
def tearDown(self):
yield self.crawler.stop()
@inlineCallbacks
def test_integration_downloader_aware_priority_queue(self):
with MockServer() as mockserver:

View File

@ -6,7 +6,6 @@ import pytest
from testfixtures import LogCapture
from twisted.internet import defer
from twisted.internet.defer import inlineCallbacks
from twisted.trial.unittest import TestCase
from scrapy.core.scheduler import BaseScheduler
from scrapy.http import Request
@ -115,8 +114,8 @@ class TestMinimalScheduler(InterfaceCheckMixin):
assert not self.scheduler.has_pending_requests()
class TestSimpleScheduler(TestCase, InterfaceCheckMixin):
def setUp(self):
class TestSimpleScheduler(InterfaceCheckMixin):
def setup_method(self):
self.scheduler = SimpleScheduler()
@inlineCallbacks
@ -145,7 +144,7 @@ class TestSimpleScheduler(TestCase, InterfaceCheckMixin):
assert close_result == "close"
class TestMinimalSchedulerCrawl(TestCase):
class TestMinimalSchedulerCrawl:
scheduler_cls = MinimalScheduler
@inlineCallbacks

View File

@ -1,6 +1,5 @@
import pytest
from twisted.internet.defer import inlineCallbacks
from twisted.trial.unittest import TestCase
from scrapy import Request, Spider, signals
from scrapy.utils.defer import deferred_f_from_coro_f, maybe_deferred_to_future
@ -21,7 +20,7 @@ class ItemSpider(Spider):
return {"index": response.meta["index"]}
class TestMain(TestCase):
class TestMain:
@deferred_f_from_coro_f
async def test_scheduler_empty(self):
crawler = get_crawler()
@ -35,17 +34,17 @@ class TestMain(TestCase):
assert len(calls) >= 1
class TestMockServer(TestCase):
class TestMockServer:
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
def setUp(self):
def setup_method(self):
self.items = []
async def _on_item_scraped(self, item):

View File

@ -13,7 +13,6 @@ from unittest import mock
import pytest
from testfixtures import LogCapture
from twisted.internet.defer import inlineCallbacks
from twisted.trial import unittest
from w3lib.url import safe_url_string
from scrapy import signals
@ -35,15 +34,9 @@ from scrapy.utils.test import get_crawler, get_reactor_settings
from tests import get_testdata, tests_datadir
class TestSpider(unittest.TestCase):
class TestSpider:
spider_class = Spider
def setUp(self):
warnings.simplefilter("always")
def tearDown(self):
warnings.resetwarnings()
def test_base_spider(self):
spider = self.spider_class("example.com")
assert spider.name == "example.com"

View File

@ -6,7 +6,6 @@ from typing import Any
import pytest
from testfixtures import LogCapture
from twisted.trial.unittest import TestCase
from scrapy import Spider, signals
from scrapy.exceptions import ScrapyDeprecationWarning
@ -21,7 +20,7 @@ ITEM_A = {"id": "a"}
ITEM_B = {"id": "b"}
class TestMain(TestCase):
class TestMain:
async def _test_spider(
self, spider: type[Spider], expected_items: list[Any] | None = None
) -> None:

View File

@ -8,7 +8,6 @@ from unittest import mock
import pytest
from testfixtures import LogCapture
from twisted.internet import defer
from twisted.trial.unittest import TestCase
from scrapy.core.spidermw import SpiderMiddlewareManager
from scrapy.exceptions import _InvalidOutput
@ -22,8 +21,8 @@ if TYPE_CHECKING:
from twisted.python.failure import Failure
class TestSpiderMiddleware(TestCase):
def setUp(self):
class TestSpiderMiddleware:
def setup_method(self):
self.request = Request("http://example.com/index.html")
self.response = Response(self.request.url, request=self.request)
self.crawler = get_crawler(Spider, {"SPIDER_MIDDLEWARES_BASE": {}})

View File

@ -5,7 +5,6 @@ import logging
import pytest
from testfixtures import LogCapture
from twisted.internet.defer import inlineCallbacks
from twisted.trial.unittest import TestCase
from scrapy.http import Request, Response
from scrapy.settings import Settings
@ -205,14 +204,14 @@ class TestHttpErrorMiddlewareHandleAll:
mw.process_spider_input(res402, spider)
class TestHttpErrorMiddlewareIntegrational(TestCase):
class TestHttpErrorMiddlewareIntegrational:
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
@inlineCallbacks

View File

@ -1,5 +1,4 @@
from testfixtures import LogCapture
from twisted.trial.unittest import TestCase
from scrapy import Request, Spider
from scrapy.utils.defer import deferred_f_from_coro_f, maybe_deferred_to_future
@ -298,16 +297,16 @@ class NotGeneratorOutputChainSpider(Spider):
# ================================================================================
class TestSpiderMiddleware(TestCase):
class TestSpiderMiddleware:
mockserver: MockServer
@classmethod
def setUpClass(cls):
def setup_class(cls):
cls.mockserver = MockServer()
cls.mockserver.__enter__()
@classmethod
def tearDownClass(cls):
def teardown_class(cls):
cls.mockserver.__exit__(None, None, None)
async def crawl_log(self, spider: type[Spider]) -> LogCapture:

View File

@ -2,7 +2,6 @@ import warnings
from asyncio import sleep
import pytest
from twisted.trial.unittest import TestCase
from scrapy import Spider, signals
from scrapy.exceptions import ScrapyDeprecationWarning
@ -106,7 +105,7 @@ class DeprecatedWrapSpiderMiddleware:
yield ITEM_C
class TestMain(TestCase):
class TestMain:
async def _test(self, spider_middlewares, spider_cls, expected_items):
actual_items = []

View File

@ -1,5 +1,3 @@
from twisted.trial.unittest import TestCase
from scrapy.http import Request
from scrapy.spidermiddlewares.start import StartSpiderMiddleware
from scrapy.spiders import Spider
@ -8,7 +6,7 @@ from scrapy.utils.misc import build_from_crawler
from scrapy.utils.test import get_crawler
class TestMiddleware(TestCase):
class TestMiddleware:
@deferred_f_from_coro_f
async def test_async(self):
crawler = get_crawler(Spider)

View File

@ -1,10 +1,8 @@
from twisted.trial import unittest
from scrapy.utils.asyncgen import as_async_generator, collect_asyncgen
from scrapy.utils.defer import deferred_f_from_coro_f
class TestAsyncgenUtils(unittest.TestCase):
class TestAsyncgenUtils:
@deferred_f_from_coro_f
async def test_as_async_generator(self):
ag = as_async_generator(range(42))

View File

@ -7,7 +7,6 @@ from unittest import mock
import pytest
from twisted.internet.defer import Deferred
from twisted.trial import unittest
from scrapy.utils.asyncgen import as_async_generator
from scrapy.utils.asyncio import (
@ -21,15 +20,14 @@ if TYPE_CHECKING:
from collections.abc import AsyncGenerator
@pytest.mark.usefixtures("reactor_pytest")
class TestAsyncio:
def test_is_asyncio_available(self):
def test_is_asyncio_available(self, reactor_pytest: str) -> None:
# the result should depend only on the pytest --reactor argument
assert is_asyncio_available() == (self.reactor_pytest != "default")
assert is_asyncio_available() == (reactor_pytest == "asyncio")
@pytest.mark.only_asyncio
class TestParallelAsyncio(unittest.TestCase):
class TestParallelAsyncio:
"""Test for scrapy.utils.asyncio.parallel_asyncio(), based on tests.test_utils_defer.TestParallelAsync."""
CONCURRENT_ITEMS = 50

View File

@ -8,7 +8,6 @@ from typing import TYPE_CHECKING, Any
import pytest
from twisted.internet.defer import Deferred, inlineCallbacks, succeed
from twisted.python.failure import Failure
from twisted.trial import unittest
from scrapy.utils.asyncgen import as_async_generator, collect_asyncgen
from scrapy.utils.defer import (
@ -29,7 +28,7 @@ if TYPE_CHECKING:
@pytest.mark.filterwarnings("ignore::scrapy.exceptions.ScrapyDeprecationWarning")
class TestMustbeDeferred(unittest.TestCase):
class TestMustbeDeferred:
@inlineCallbacks
def test_success_function(self) -> Generator[Deferred[Any], Any, None]:
steps: list[int] = []
@ -87,7 +86,7 @@ def eb1(failure, arg1, arg2):
return f"(eb1 {failure.value.__class__.__name__} {arg1} {arg2})"
class TestDeferUtils(unittest.TestCase):
class TestDeferUtils:
@inlineCallbacks
def test_process_chain(self):
x = yield process_chain([cb1, cb2, cb3], "res", "v1", "v2")
@ -131,7 +130,7 @@ class TestIterErrback:
assert isinstance(errors[0].value, ZeroDivisionError)
class TestAiterErrback(unittest.TestCase):
class TestAiterErrback:
@deferred_f_from_coro_f
async def test_aiter_errback_good(self):
async def itergood() -> AsyncGenerator[int, None]:
@ -158,7 +157,7 @@ class TestAiterErrback(unittest.TestCase):
assert isinstance(errors[0].value, ZeroDivisionError)
class TestAsyncDefTestsuite(unittest.TestCase):
class TestAsyncDefTestsuite:
@deferred_f_from_coro_f
async def test_deferred_f_from_coro_f(self):
pass
@ -173,7 +172,7 @@ class TestAsyncDefTestsuite(unittest.TestCase):
raise RuntimeError("This is expected to be raised")
class TestParallelAsync(unittest.TestCase):
class TestParallelAsync:
"""This tests _AsyncCooperatorAdapter by testing parallel_async which is its only usage.
parallel_async is called with the results of a callback (so an iterable of items, requests and None,
@ -283,7 +282,7 @@ class TestParallelAsync(unittest.TestCase):
assert max_parallel_count[0] <= self.CONCURRENT_ITEMS, max_parallel_count[0]
class TestDeferredFromCoro(unittest.TestCase):
class TestDeferredFromCoro:
def test_deferred(self):
d = Deferred()
result = deferred_from_coro(d)
@ -327,7 +326,7 @@ class TestDeferredFromCoro(unittest.TestCase):
assert future_result == 42
class TestDeferredFFromCoroF(unittest.TestCase):
class TestDeferredFFromCoroF:
@inlineCallbacks
def _assert_result(
self, c_f: Callable[[], Awaitable[int]]
@ -364,7 +363,7 @@ class TestDeferredFFromCoroF(unittest.TestCase):
@pytest.mark.only_asyncio
class TestDeferredToFuture(unittest.TestCase):
class TestDeferredToFuture:
@deferred_f_from_coro_f
async def test_deferred(self):
d = Deferred()
@ -399,7 +398,7 @@ class TestDeferredToFuture(unittest.TestCase):
@pytest.mark.only_asyncio
class TestMaybeDeferredToFutureAsyncio(unittest.TestCase):
class TestMaybeDeferredToFutureAsyncio:
@deferred_f_from_coro_f
async def test_deferred(self):
d = Deferred()

View File

@ -7,7 +7,6 @@ import sys
from typing import TYPE_CHECKING, TypeVar
import pytest
from twisted.trial import unittest
from scrapy.utils.asyncgen import as_async_generator, collect_asyncgen
from scrapy.utils.defer import aiter_errback, deferred_f_from_coro_f
@ -41,7 +40,7 @@ def test_mutablechain():
assert list(m) == list(range(2, 13))
class TestMutableAsyncChain(unittest.TestCase):
class TestMutableAsyncChain:
@staticmethod
async def g1():
for i in range(3):

View File

@ -2,7 +2,6 @@ import asyncio
import warnings
import pytest
from twisted.trial.unittest import TestCase
from scrapy.utils.defer import deferred_f_from_coro_f
from scrapy.utils.reactor import (
@ -13,11 +12,10 @@ from scrapy.utils.reactor import (
)
@pytest.mark.usefixtures("reactor_pytest")
class TestAsyncio(TestCase):
def test_is_asyncio_reactor_installed(self):
class TestAsyncio:
def test_is_asyncio_reactor_installed(self, reactor_pytest: str) -> None:
# the result should depend only on the pytest --reactor argument
assert is_asyncio_reactor_installed() == (self.reactor_pytest != "default")
assert is_asyncio_reactor_installed() == (reactor_pytest == "asyncio")
def test_install_asyncio_reactor(self):
from twisted.internet import reactor as original_reactor

View File

@ -6,7 +6,6 @@ from testfixtures import LogCapture
from twisted.internet import defer
from twisted.internet.defer import inlineCallbacks
from twisted.python.failure import Failure
from twisted.trial import unittest
from scrapy.utils.defer import deferred_from_coro
from scrapy.utils.signal import (
@ -17,7 +16,7 @@ from scrapy.utils.signal import (
from scrapy.utils.test import get_from_asyncio_queue
class TestSendCatchLog(unittest.TestCase):
class TestSendCatchLog:
@inlineCallbacks
def test_send_catch_log(self):
test_signal = object()
@ -75,7 +74,6 @@ class TestSendCatchLogDeferred2(TestSendCatchLogDeferred):
return d
@pytest.mark.usefixtures("reactor_pytest")
class TestSendCatchLogDeferredAsyncDef(TestSendCatchLogDeferred):
async def ok_handler(self, arg, handlers_called):
handlers_called.add(self.ok_handler)
@ -109,7 +107,6 @@ class TestSendCatchLogAsync2(TestSendCatchLogAsync):
return d
@pytest.mark.usefixtures("reactor_pytest")
class TestSendCatchLogAsyncAsyncDef(TestSendCatchLogAsync):
async def ok_handler(self, arg, handlers_called):
handlers_called.add(self.ok_handler)

View File

@ -1,22 +1,18 @@
"""
from twisted.internet import defer
Tests borrowed from the twisted.web.client tests.
"""
from __future__ import annotations
import shutil
from pathlib import Path
from tempfile import mkdtemp
from urllib.parse import urlparse
import OpenSSL.SSL
import pytest
from pytest_twisted import async_yield_fixture
from twisted.internet import defer
from twisted.internet.defer import inlineCallbacks
from twisted.internet.testing import StringTransport
from twisted.protocols.policies import WrappingFactory
from twisted.trial import unittest
from twisted.web import resource, server, static, util
from twisted.web.client import _makeGetterFactory
@ -200,16 +196,16 @@ class EncodingResource(resource.Resource):
@pytest.mark.filterwarnings("ignore::scrapy.exceptions.ScrapyDeprecationWarning")
class TestWebClient(unittest.TestCase):
class TestWebClient:
def _listen(self, site):
from twisted.internet import reactor
return reactor.listenTCP(0, site, interface="127.0.0.1")
def setUp(self):
self.tmpname = Path(mkdtemp())
(self.tmpname / "file").write_bytes(b"0123456789")
r = static.File(str(self.tmpname))
@pytest.fixture
def wrapper(self, tmp_path):
(tmp_path / "file").write_bytes(b"0123456789")
r = static.File(str(tmp_path))
r.putChild(b"redirect", util.Redirect(b"/file"))
r.putChild(b"wait", ForeverTakingResource())
r.putChild(b"error", ErrorResource())
@ -218,45 +214,47 @@ class TestWebClient(unittest.TestCase):
r.putChild(b"payload", PayloadResource())
r.putChild(b"broken", BrokenDownloadResource())
r.putChild(b"encoding", EncodingResource())
self.site = server.Site(r, timeout=None)
self.wrapper = WrappingFactory(self.site)
self.port = self._listen(self.wrapper)
self.portno = self.port.getHost().port
site = server.Site(r, timeout=None)
return WrappingFactory(site)
@async_yield_fixture
async def server_port(self, wrapper):
port = self._listen(wrapper)
yield port.getHost().port
await port.stopListening()
@pytest.fixture
def server_url(self, server_port):
return f"http://127.0.0.1:{server_port}/"
@inlineCallbacks
def tearDown(self):
yield self.port.stopListening()
shutil.rmtree(self.tmpname)
def getURL(self, path):
return f"http://127.0.0.1:{self.portno}/{path}"
@inlineCallbacks
def testPayload(self):
def testPayload(self, server_url):
s = "0123456789" * 10
body = yield getPage(self.getURL("payload"), body=s)
body = yield getPage(server_url + "payload", body=s)
assert body == to_bytes(s)
@inlineCallbacks
def testHostHeader(self):
def testHostHeader(self, server_port, server_url):
# if we pass Host header explicitly, it should be used, otherwise
# it should extract from url
body = yield getPage(self.getURL("host"))
assert body == to_bytes(f"127.0.0.1:{self.portno}")
body = yield getPage(self.getURL("host"), headers={"Host": "www.example.com"})
body = yield getPage(server_url + "host")
assert body == to_bytes(f"127.0.0.1:{server_port}")
body = yield getPage(server_url + "host", headers={"Host": "www.example.com"})
assert body == to_bytes("www.example.com")
@inlineCallbacks
def test_getPage(self):
def test_getPage(self, server_url):
"""
L{client.getPage} returns a L{Deferred} which is called back with
the body of the response if the default method B{GET} is used.
"""
body = yield getPage(self.getURL("file"))
body = yield getPage(server_url + "file")
assert body == b"0123456789"
@inlineCallbacks
def test_getPageHead(self):
def test_getPageHead(self, server_url):
"""
L{client.getPage} returns a L{Deferred} which is called back with
the empty string if the method is C{HEAD} and there is a successful
@ -264,7 +262,7 @@ class TestWebClient(unittest.TestCase):
"""
def _getPage(method):
return getPage(self.getURL("file"), method=method)
return getPage(server_url + "file", method=method)
body = yield _getPage("head")
assert body == b""
@ -272,42 +270,42 @@ class TestWebClient(unittest.TestCase):
assert body == b""
@inlineCallbacks
def test_timeoutNotTriggering(self):
def test_timeoutNotTriggering(self, server_port, server_url):
"""
When a non-zero timeout is passed to L{getPage} and the page is
retrieved before the timeout period elapses, the L{Deferred} is
called back with the contents of the page.
"""
body = yield getPage(self.getURL("host"), timeout=100)
assert body == to_bytes(f"127.0.0.1:{self.portno}")
body = yield getPage(server_url + "host", timeout=100)
assert body == to_bytes(f"127.0.0.1:{server_port}")
@inlineCallbacks
def test_timeoutTriggering(self):
def test_timeoutTriggering(self, wrapper, server_url):
"""
When a non-zero timeout is passed to L{getPage} and that many
seconds elapse before the server responds to the request. the
L{Deferred} is errbacked with a L{error.TimeoutError}.
"""
with pytest.raises(defer.TimeoutError):
yield getPage(self.getURL("wait"), timeout=0.000001)
yield getPage(server_url + "wait", timeout=0.000001)
# Clean up the server which is hanging around not doing
# anything.
connected = list(self.wrapper.protocols.keys())
connected = list(wrapper.protocols.keys())
# There might be nothing here if the server managed to already see
# that the connection was lost.
if connected:
connected[0].transport.loseConnection()
@inlineCallbacks
def testNotFound(self):
body = yield getPage(self.getURL("notsuchfile"))
def testNotFound(self, server_url):
body = yield getPage(server_url + "notsuchfile")
assert b"404 - No Such Resource" in body
@inlineCallbacks
def testFactoryInfo(self):
def testFactoryInfo(self, server_url):
from twisted.internet import reactor
url = self.getURL("file")
url = server_url + "file"
parsed = urlparse(url)
factory = client.ScrapyHTTPClientFactory(Request(url))
reactor.connectTCP(parsed.hostname, parsed.port, factory)
@ -318,8 +316,8 @@ class TestWebClient(unittest.TestCase):
assert factory.response_headers[b"content-length"] == b"10"
@inlineCallbacks
def testRedirect(self):
body = yield getPage(self.getURL("redirect"))
def testRedirect(self, server_url):
body = yield getPage(server_url + "redirect")
assert (
body
== b'\n<html>\n <head>\n <meta http-equiv="refresh" content="0;URL=/file">\n'
@ -328,12 +326,12 @@ class TestWebClient(unittest.TestCase):
)
@inlineCallbacks
def test_encoding(self):
def test_encoding(self, server_url):
"""Test that non-standart body encoding matches
Content-Encoding header"""
original_body = b"\xd0\x81\xd1\x8e\xd0\xaf"
response = yield getPage(
self.getURL("encoding"), body=original_body, response_transform=lambda r: r
server_url + "encoding", body=original_body, response_transform=lambda r: r
)
content_encoding = to_unicode(response.headers[b"Content-Encoding"])
assert content_encoding == EncodingResource.out_encoding
@ -343,9 +341,9 @@ class TestWebClient(unittest.TestCase):
@pytest.mark.filterwarnings("ignore::scrapy.exceptions.ScrapyDeprecationWarning")
class TestWebClientSSL(TestContextFactoryBase):
@inlineCallbacks
def testPayload(self):
def testPayload(self, server_url):
s = "0123456789" * 10
body = yield getPage(self.getURL("payload"), body=s)
body = yield getPage(server_url + "payload", body=s)
assert body == to_bytes(s)
@ -355,19 +353,19 @@ class TestWebClientCustomCiphersSSL(TestWebClientSSL):
context_factory = ssl_context_factory(cipher_string=custom_ciphers)
@inlineCallbacks
def testPayload(self):
def testPayload(self, server_url):
s = "0123456789" * 10
crawler = get_crawler(
settings_dict={"DOWNLOADER_CLIENT_TLS_CIPHERS": self.custom_ciphers}
)
client_context_factory = build_from_crawler(ScrapyClientContextFactory, crawler)
body = yield getPage(
self.getURL("payload"), body=s, contextFactory=client_context_factory
server_url + "payload", body=s, contextFactory=client_context_factory
)
assert body == to_bytes(s)
@inlineCallbacks
def testPayloadDisabledCipher(self):
def testPayloadDisabledCipher(self, server_url):
s = "0123456789" * 10
crawler = get_crawler(
settings_dict={
@ -377,5 +375,5 @@ class TestWebClientCustomCiphersSSL(TestWebClientSSL):
client_context_factory = build_from_crawler(ScrapyClientContextFactory, crawler)
with pytest.raises(OpenSSL.SSL.Error):
yield getPage(
self.getURL("payload"), body=s, contextFactory=client_context_factory
server_url + "payload", body=s, contextFactory=client_context_factory
)

View File

@ -19,6 +19,7 @@ deps =
pytest-xdist
sybil >= 1.3.0 # https://github.com/cjw296/sybil/issues/20#issuecomment-605433422
testfixtures
pytest-twisted >= 1.14.3
[testenv]
deps =