From 6b2997af9000bbf9114aca681b62713c30cc628e Mon Sep 17 00:00:00 2001 From: Andrey Rakhmatullin Date: Sun, 6 Jul 2025 21:27:17 +0500 Subject: [PATCH] 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. --- conftest.py | 37 +- pyproject.toml | 3 + scrapy/utils/reactor.py | 7 +- tests/test_addons.py | 3 +- tests/test_closespider.py | 7 +- tests/test_contracts.py | 5 +- tests/test_core_downloader.py | 65 +-- tests/test_crawl.py | 13 +- tests/test_crawler.py | 15 +- .../test_downloader_handler_twisted_http10.py | 15 +- .../test_downloader_handler_twisted_http2.py | 111 ++-- tests/test_downloader_handlers.py | 262 +++++---- tests/test_downloader_handlers_http_base.py | 534 ++++++++++-------- tests/test_downloadermiddleware.py | 131 +++-- tests/test_downloadermiddleware_robotstxt.py | 11 +- tests/test_downloaderslotssettings.py | 9 +- tests/test_engine.py | 3 +- tests/test_engine_loop.py | 9 +- tests/test_extension_telnet.py | 3 +- tests/test_feedexport.py | 29 +- tests/test_http2_client_protocol.py | 415 ++++++++------ tests/test_http_request.py | 6 - tests/test_logformatter.py | 9 +- tests/test_pipeline_crawl.py | 11 +- tests/test_pipeline_files.py | 13 +- tests/test_pipeline_media.py | 13 +- tests/test_pipelines.py | 7 +- tests/test_proxy_connect.py | 11 +- tests/test_request_attribute_binding.py | 7 +- tests/test_request_cb_kwargs.py | 7 +- tests/test_request_left.py | 7 +- tests/test_scheduler.py | 9 +- tests/test_scheduler_base.py | 7 +- tests/test_signals.py | 11 +- tests/test_spider.py | 9 +- tests/test_spider_start.py | 3 +- tests/test_spidermiddleware.py | 5 +- tests/test_spidermiddleware_httperror.py | 7 +- tests/test_spidermiddleware_output_chain.py | 7 +- tests/test_spidermiddleware_process_start.py | 3 +- tests/test_spidermiddleware_start.py | 4 +- tests/test_utils_asyncgen.py | 4 +- tests/test_utils_asyncio.py | 8 +- tests/test_utils_defer.py | 19 +- tests/test_utils_python.py | 3 +- tests/test_utils_reactor.py | 8 +- tests/test_utils_signal.py | 5 +- tests/test_webclient.py | 102 ++-- tox.ini | 1 + 49 files changed, 1054 insertions(+), 939 deletions(-) diff --git a/conftest.py b/conftest.py index 02697603d..4cfacc2a2 100644 --- a/conftest.py +++ b/conftest.py @@ -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 diff --git a/pyproject.toml b/pyproject.toml index 315857eda..9db636a08 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -218,6 +218,9 @@ disable = [ ] [tool.pytest.ini_options] +addopts = [ + "--reactor=asyncio", +] xfail_strict = true python_files = ["test_*.py", "test_*/__init__.py"] markers = [ diff --git a/scrapy/utils/reactor.py b/scrapy/utils/reactor.py index 373ad652c..132f88c74 100644 --- a/scrapy/utils/reactor.py +++ b/scrapy/utils/reactor.py @@ -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: diff --git a/tests/test_addons.py b/tests/test_addons.py index b4294c815..0383fa627 100644 --- a/tests/test_addons.py +++ b/tests/test_addons.py @@ -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}, diff --git a/tests/test_closespider.py b/tests/test_closespider.py index c6ec690a1..563ecbe92 100644 --- a/tests/test_closespider.py +++ b/tests/test_closespider.py @@ -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 diff --git a/tests/test_contracts.py b/tests/test_contracts.py index ad3efa042..fc3cd9df0 100644 --- a/tests/test_contracts.py +++ b/tests/test_contracts.py @@ -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) diff --git a/tests/test_core_downloader.py b/tests/test_core_downloader.py index 0a674d638..ca15c560a 100644 --- a/tests/test_core_downloader.py +++ b/tests/test_core_downloader.py @@ -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) diff --git a/tests/test_crawl.py b/tests/test_crawl.py index ef998ab4c..877b23bef 100644 --- a/tests/test_crawl.py +++ b/tests/test_crawl.py @@ -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( diff --git a/tests/test_crawler.py b/tests/test_crawler.py index e44bf1a46..bc182d2f5 100644 --- a/tests/test_crawler.py +++ b/tests/test_crawler.py @@ -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. """ diff --git a/tests/test_downloader_handler_twisted_http10.py b/tests/test_downloader_handler_twisted_http10.py index 4b700be6b..ddb3250db 100644 --- a/tests/test_downloader_handler_twisted_http10.py +++ b/tests/test_downloader_handler_twisted_http10.py @@ -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" diff --git a/tests/test_downloader_handler_twisted_http2.py b/tests/test_downloader_handler_twisted_http2.py index 2dc6817a5..a76cf9dfc 100644 --- a/tests/test_downloader_handler_twisted_http2.py +++ b/tests/test_downloader_handler_twisted_http2.py @@ -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 + ) ) diff --git a/tests/test_downloader_handlers.py b/tests/test_downloader_handlers.py index e73bfbccf..eba14c0c3 100644 --- a/tests/test_downloader_handlers.py +++ b/tests/test_downloader_handlers.py @@ -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"\n."), ) - 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 diff --git a/tests/test_downloader_handlers_http_base.py b/tests/test_downloader_handlers_http_base.py index 1cdbbc831..35f5d483e 100644 --- a/tests/test_downloader_handlers_http_base.py +++ b/tests/test_downloader_handlers_http_base.py @@ -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"\n.", 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"\n.", - 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 diff --git a/tests/test_downloadermiddleware.py b/tests/test_downloadermiddleware.py index d12baf1ad..cfab0966a 100644 --- a/tests/test_downloadermiddleware.py +++ b/tests/test_downloadermiddleware.py @@ -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 diff --git a/tests/test_downloadermiddleware_robotstxt.py b/tests/test_downloadermiddleware_robotstxt.py index dd5d47cab..12e43800b 100644 --- a/tests/test_downloadermiddleware_robotstxt.py +++ b/tests/test_downloadermiddleware_robotstxt.py @@ -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" ) diff --git a/tests/test_downloaderslotssettings.py b/tests/test_downloaderslotssettings.py index ddac95edf..0d9500464 100644 --- a/tests/test_downloaderslotssettings.py +++ b/tests/test_downloaderslotssettings.py @@ -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 diff --git a/tests/test_engine.py b/tests/test_engine.py index 0f8cbc3b5..ecb615f61 100644 --- a/tests/test_engine.py +++ b/tests/test_engine.py @@ -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 = [ diff --git a/tests/test_engine_loop.py b/tests/test_engine_loop.py index d4c7b8698..49a800fe2 100644 --- a/tests/test_engine_loop.py +++ b/tests/test_engine_loop.py @@ -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): diff --git a/tests/test_extension_telnet.py b/tests/test_extension_telnet.py index f9e54cb28..6b4ad450f 100644 --- a/tests/test_extension_telnet.py +++ b/tests/test_extension_telnet.py @@ -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) diff --git a/tests/test_feedexport.py b/tests/test_feedexport.py index 9085af18e..309466b90 100644 --- a/tests/test_feedexport.py +++ b/tests/test_feedexport.py @@ -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]], diff --git a/tests/test_http2_client_protocol.py b/tests/test_http2_client_protocol.py index 81c507ea1..77d328333 100644 --- a/tests/test_http2_client_protocol.py +++ b/tests/test_http2_client_protocol.py @@ -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] diff --git a/tests/test_http_request.py b/tests/test_http_request.py index 90a693bfa..3a62bf716 100644 --- a/tests/test_http_request.py +++ b/tests/test_http_request.py @@ -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"" diff --git a/tests/test_logformatter.py b/tests/test_logformatter.py index 682428954..4cc48ca51 100644 --- a/tests/test_logformatter.py +++ b/tests/test_logformatter.py @@ -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": { diff --git a/tests/test_pipeline_crawl.py b/tests/test_pipeline_crawl.py index cf827e481..5f758a763 100644 --- a/tests/test_pipeline_crawl.py +++ b/tests/test_pipeline_crawl.py @@ -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 = [] diff --git a/tests/test_pipeline_files.py b/tests/test_pipeline_files.py index 3d84c7e71..38d3b8ce5 100644 --- a/tests/test_pipeline_files.py +++ b/tests/test_pipeline_files.py @@ -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" diff --git a/tests/test_pipeline_media.py b/tests/test_pipeline_media.py index 5364b09a8..b309c5d8e 100644 --- a/tests/test_pipeline_media.py +++ b/tests/test_pipeline_media.py @@ -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) diff --git a/tests/test_pipelines.py b/tests/test_pipelines.py index ea85877bf..4df9495a5 100644 --- a/tests/test_pipelines.py +++ b/tests/test_pipelines.py @@ -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): diff --git a/tests/test_proxy_connect.py b/tests/test_proxy_connect.py index 8875dc04f..400dfa4ba 100644 --- a/tests/test_proxy_connect.py +++ b/tests/test_proxy_connect.py @@ -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 diff --git a/tests/test_request_attribute_binding.py b/tests/test_request_attribute_binding.py index 9318ee87e..c0606ac35 100644 --- a/tests/test_request_attribute_binding.py +++ b/tests/test_request_attribute_binding.py @@ -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 diff --git a/tests/test_request_cb_kwargs.py b/tests/test_request_cb_kwargs.py index 1714bd4db..34a07a3a1 100644 --- a/tests/test_request_cb_kwargs.py +++ b/tests/test_request_cb_kwargs.py @@ -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 diff --git a/tests/test_request_left.py b/tests/test_request_left.py index 12ef42610..451125932 100644 --- a/tests/test_request_left.py +++ b/tests/test_request_left.py @@ -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 diff --git a/tests/test_scheduler.py b/tests/test_scheduler.py index d31ba48e5..9aed270e3 100644 --- a/tests/test_scheduler.py +++ b/tests/test_scheduler.py @@ -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: diff --git a/tests/test_scheduler_base.py b/tests/test_scheduler_base.py index 26482fc8d..f85a754c2 100644 --- a/tests/test_scheduler_base.py +++ b/tests/test_scheduler_base.py @@ -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 diff --git a/tests/test_signals.py b/tests/test_signals.py index b20a949e8..dcbd0fb35 100644 --- a/tests/test_signals.py +++ b/tests/test_signals.py @@ -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): diff --git a/tests/test_spider.py b/tests/test_spider.py index e2843eca4..4a532ed78 100644 --- a/tests/test_spider.py +++ b/tests/test_spider.py @@ -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" diff --git a/tests/test_spider_start.py b/tests/test_spider_start.py index 74bb5c87c..d4eca85b8 100644 --- a/tests/test_spider_start.py +++ b/tests/test_spider_start.py @@ -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: diff --git a/tests/test_spidermiddleware.py b/tests/test_spidermiddleware.py index 8d9f3e83d..b0f69c1ff 100644 --- a/tests/test_spidermiddleware.py +++ b/tests/test_spidermiddleware.py @@ -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": {}}) diff --git a/tests/test_spidermiddleware_httperror.py b/tests/test_spidermiddleware_httperror.py index c1511a9a4..5289335e3 100644 --- a/tests/test_spidermiddleware_httperror.py +++ b/tests/test_spidermiddleware_httperror.py @@ -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 diff --git a/tests/test_spidermiddleware_output_chain.py b/tests/test_spidermiddleware_output_chain.py index 60464d696..d8af25b3a 100644 --- a/tests/test_spidermiddleware_output_chain.py +++ b/tests/test_spidermiddleware_output_chain.py @@ -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: diff --git a/tests/test_spidermiddleware_process_start.py b/tests/test_spidermiddleware_process_start.py index e1c8b5fec..a525f991d 100644 --- a/tests/test_spidermiddleware_process_start.py +++ b/tests/test_spidermiddleware_process_start.py @@ -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 = [] diff --git a/tests/test_spidermiddleware_start.py b/tests/test_spidermiddleware_start.py index 295b10ea8..c2efe47c9 100644 --- a/tests/test_spidermiddleware_start.py +++ b/tests/test_spidermiddleware_start.py @@ -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) diff --git a/tests/test_utils_asyncgen.py b/tests/test_utils_asyncgen.py index 9b5a25b3a..b2d2b04c7 100644 --- a/tests/test_utils_asyncgen.py +++ b/tests/test_utils_asyncgen.py @@ -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)) diff --git a/tests/test_utils_asyncio.py b/tests/test_utils_asyncio.py index a6e52eb26..b6b77c9e3 100644 --- a/tests/test_utils_asyncio.py +++ b/tests/test_utils_asyncio.py @@ -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 diff --git a/tests/test_utils_defer.py b/tests/test_utils_defer.py index b31d8c3d9..ecb269709 100644 --- a/tests/test_utils_defer.py +++ b/tests/test_utils_defer.py @@ -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() diff --git a/tests/test_utils_python.py b/tests/test_utils_python.py index c933e0ac9..8e1105b09 100644 --- a/tests/test_utils_python.py +++ b/tests/test_utils_python.py @@ -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): diff --git a/tests/test_utils_reactor.py b/tests/test_utils_reactor.py index eb00ab193..6cbdcfccc 100644 --- a/tests/test_utils_reactor.py +++ b/tests/test_utils_reactor.py @@ -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 diff --git a/tests/test_utils_signal.py b/tests/test_utils_signal.py index 79bac8bc5..d90d3ed95 100644 --- a/tests/test_utils_signal.py +++ b/tests/test_utils_signal.py @@ -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) diff --git a/tests/test_webclient.py b/tests/test_webclient.py index dd6d7939e..d441a03a9 100644 --- a/tests/test_webclient.py +++ b/tests/test_webclient.py @@ -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\n \n \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 ) diff --git a/tox.ini b/tox.ini index 1c1d7c0ec..db6add080 100644 --- a/tox.ini +++ b/tox.ini @@ -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 =