From 044c3f69edd1bf926408649361ada7f2146db04e Mon Sep 17 00:00:00 2001 From: Mehraz Hossain Rumman <59512321+MehrazRumman@users.noreply.github.com> Date: Mon, 10 Mar 2025 01:18:57 +0600 Subject: [PATCH 1/5] Deprecate InitSpider (#6714) --- scrapy/spiders/init.py | 17 ++++++++++++++++- tests/test_spider.py | 1 + 2 files changed, 17 insertions(+), 1 deletion(-) diff --git a/scrapy/spiders/init.py b/scrapy/spiders/init.py index 4ec2919f7..a7dba989e 100644 --- a/scrapy/spiders/init.py +++ b/scrapy/spiders/init.py @@ -1,9 +1,11 @@ from __future__ import annotations +import warnings from collections.abc import Iterable from typing import TYPE_CHECKING, Any, cast from scrapy import Request +from scrapy.exceptions import ScrapyDeprecationWarning from scrapy.spiders import Spider from scrapy.utils.spider import iterate_spider_output @@ -12,7 +14,20 @@ if TYPE_CHECKING: class InitSpider(Spider): - """Base Spider with initialization facilities""" + """Base Spider with initialization facilities + + .. warning:: This class is deprecated. Copy its code into your project if needed. + It will be removed in a future Scrapy version. + """ + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + warnings.warn( + "InitSpider is deprecated. Copy its code from Scrapy's source if needed. " + "Will be removed in a future version.", + ScrapyDeprecationWarning, + stacklevel=2, + ) def start_requests(self) -> Iterable[Request]: self._postinit_reqs: Iterable[Request] = super().start_requests() diff --git a/tests/test_spider.py b/tests/test_spider.py index 05f1c59d0..4e8330c06 100644 --- a/tests/test_spider.py +++ b/tests/test_spider.py @@ -144,6 +144,7 @@ class TestSpider(unittest.TestCase): mock_logger.log.assert_called_once_with("INFO", "test log msg") +@pytest.mark.filterwarnings("ignore::scrapy.exceptions.ScrapyDeprecationWarning") class TestInitSpider(TestSpider): spider_class = InitSpider From 02ed71d8877d1f3f270a9085c3cdb7fc7e917b8a Mon Sep 17 00:00:00 2001 From: Andrey Rakhmatullin Date: Sun, 9 Mar 2025 23:20:24 +0400 Subject: [PATCH 2/5] Converting tests to plain asserts, part 6. (#6709) --- tests/test_item.py | 97 ++-- tests/test_link.py | 12 +- tests/test_linkextractors.py | 937 +++++++++++++++------------------- tests/test_pipeline_crawl.py | 46 +- tests/test_pipeline_files.py | 213 ++++---- tests/test_pipeline_images.py | 204 ++++---- tests/test_pipeline_media.py | 149 +++--- 7 files changed, 730 insertions(+), 928 deletions(-) diff --git a/tests/test_item.py b/tests/test_item.py index 47c5c3db6..bf51eb398 100644 --- a/tests/test_item.py +++ b/tests/test_item.py @@ -1,4 +1,3 @@ -import unittest from abc import ABCMeta from unittest import mock @@ -7,9 +6,9 @@ import pytest from scrapy.item import Field, Item, ItemMeta -class ItemTest(unittest.TestCase): +class TestItem: def assertSortedEqual(self, first, second, msg=None): - return self.assertEqual(sorted(first), sorted(second), msg) + assert sorted(first) == sorted(second), msg def test_simple(self): class TestItem(Item): @@ -17,7 +16,7 @@ class ItemTest(unittest.TestCase): i = TestItem() i["name"] = "name" - self.assertEqual(i["name"], "name") + assert i["name"] == "name" def test_init(self): class TestItem(Item): @@ -28,13 +27,13 @@ class ItemTest(unittest.TestCase): i["name"] i2 = TestItem(name="john doe") - self.assertEqual(i2["name"], "john doe") + assert i2["name"] == "john doe" i3 = TestItem({"name": "john doe"}) - self.assertEqual(i3["name"], "john doe") + assert i3["name"] == "john doe" i4 = TestItem(i3) - self.assertEqual(i4["name"], "john doe") + assert i4["name"] == "john doe" with pytest.raises(KeyError): TestItem({"name": "john doe", "other": "foo"}) @@ -59,11 +58,11 @@ class ItemTest(unittest.TestCase): i["number"] = 123 itemrepr = repr(i) - self.assertEqual(itemrepr, "{'name': 'John Doe', 'number': 123}") + assert itemrepr == "{'name': 'John Doe', 'number': 123}" i2 = eval(itemrepr) # pylint: disable=eval-used - self.assertEqual(i2["name"], "John Doe") - self.assertEqual(i2["number"], 123) + assert i2["name"] == "John Doe" + assert i2["number"] == 123 def test_private_attr(self): class TestItem(Item): @@ -71,7 +70,7 @@ class ItemTest(unittest.TestCase): i = TestItem() i._private = "test" - self.assertEqual(i._private, "test") + assert i._private == "test" def test_raise_getattr(self): class TestItem(Item): @@ -103,9 +102,9 @@ class ItemTest(unittest.TestCase): with pytest.raises(KeyError): i.get_name() i["name"] = "lala" - self.assertEqual(i.get_name(), "lala") + assert i.get_name() == "lala" i.change_name("other") - self.assertEqual(i.get_name(), "other") + assert i.get_name() == "other" def test_metaclass(self): class TestItem(Item): @@ -115,8 +114,8 @@ class ItemTest(unittest.TestCase): i = TestItem() i["name"] = "John" - self.assertEqual(list(i.keys()), ["name"]) - self.assertEqual(list(i.values()), ["John"]) + assert list(i.keys()) == ["name"] + assert list(i.values()) == ["John"] i["keys"] = "Keys" i["values"] = "Values" @@ -142,8 +141,8 @@ class ItemTest(unittest.TestCase): i = TestItem() i["keys"] = 3 - self.assertEqual(list(i.keys()), ["keys"]) - self.assertEqual(list(i.values()), [3]) + assert list(i.keys()) == ["keys"] + assert list(i.values()) == [3] def test_metaclass_multiple_inheritance_simple(self): class A(Item): @@ -161,17 +160,17 @@ class ItemTest(unittest.TestCase): pass item = D(save="X", load="Y") - self.assertEqual(item["save"], "X") - self.assertEqual(item["load"], "Y") - self.assertEqual(D.fields, {"load": {"default": "A"}, "save": {"default": "A"}}) + assert item["save"] == "X" + assert item["load"] == "Y" + assert D.fields == {"load": {"default": "A"}, "save": {"default": "A"}} # D class inverted class E(C, B): pass - self.assertEqual(E(save="X")["save"], "X") - self.assertEqual(E(load="X")["load"], "X") - self.assertEqual(E.fields, {"load": {"default": "C"}, "save": {"default": "C"}}) + assert E(save="X")["save"] == "X" + assert E(load="X")["load"] == "X" + assert E.fields == {"load": {"default": "C"}, "save": {"default": "C"}} def test_metaclass_multiple_inheritance_diamond(self): class A(Item): @@ -190,31 +189,25 @@ class ItemTest(unittest.TestCase): fields = {"update": Field(default="D")} load = Field(default="D") - self.assertEqual(D(save="X")["save"], "X") - self.assertEqual(D(load="X")["load"], "X") - self.assertEqual( - D.fields, - { - "save": {"default": "C"}, - "load": {"default": "D"}, - "update": {"default": "D"}, - }, - ) + assert D(save="X")["save"] == "X" + assert D(load="X")["load"] == "X" + assert D.fields == { + "save": {"default": "C"}, + "load": {"default": "D"}, + "update": {"default": "D"}, + } # D class inverted class E(C, B): load = Field(default="E") - self.assertEqual(E(save="X")["save"], "X") - self.assertEqual(E(load="X")["load"], "X") - self.assertEqual( - E.fields, - { - "save": {"default": "C"}, - "load": {"default": "E"}, - "update": {"default": "C"}, - }, - ) + assert E(save="X")["save"] == "X" + assert E(load="X")["load"] == "X" + assert E.fields == { + "save": {"default": "C"}, + "load": {"default": "E"}, + "update": {"default": "C"}, + } def test_metaclass_multiple_inheritance_without_metaclass(self): class A(Item): @@ -234,8 +227,8 @@ class ItemTest(unittest.TestCase): with pytest.raises(KeyError): D(not_allowed="value") - self.assertEqual(D(save="X")["save"], "X") - self.assertEqual(D.fields, {"save": {"default": "A"}, "load": {"default": "A"}}) + assert D(save="X")["save"] == "X" + assert D.fields == {"save": {"default": "A"}, "load": {"default": "A"}} # D class inverted class E(C, B): @@ -243,8 +236,8 @@ class ItemTest(unittest.TestCase): with pytest.raises(KeyError): E(not_allowed="value") - self.assertEqual(E(save="X")["save"], "X") - self.assertEqual(E.fields, {"save": {"default": "A"}, "load": {"default": "A"}}) + assert E(save="X")["save"] == "X" + assert E.fields == {"save": {"default": "A"}, "load": {"default": "A"}} def test_to_dict(self): class TestItem(Item): @@ -252,7 +245,7 @@ class ItemTest(unittest.TestCase): i = TestItem() i["name"] = "John" - self.assertEqual(dict(i), {"name": "John"}) + assert dict(i) == {"name": "John"} def test_copy(self): class TestItem(Item): @@ -260,9 +253,9 @@ class ItemTest(unittest.TestCase): item = TestItem({"name": "lower"}) copied_item = item.copy() - self.assertNotEqual(id(item), id(copied_item)) + assert id(item) != id(copied_item) copied_item["name"] = copied_item["name"].upper() - self.assertNotEqual(item["name"], copied_item["name"]) + assert item["name"] != copied_item["name"] def test_deepcopy(self): class TestItem(Item): @@ -274,7 +267,7 @@ class ItemTest(unittest.TestCase): assert item["tags"] != copied_item["tags"] -class ItemMetaTest(unittest.TestCase): +class TestItemMeta: def test_new_method_propagates_classcell(self): new_mock = mock.Mock(side_effect=ABCMeta.__new__) base = ItemMeta.__bases__[0] @@ -297,7 +290,7 @@ class ItemMetaTest(unittest.TestCase): assert "__classcell__" in attrs -class ItemMetaClassCellRegression(unittest.TestCase): +class TestItemMetaClassCellRegression: def test_item_meta_classcell_regression(self): class MyItem(Item, metaclass=ItemMeta): def __init__(self, *args, **kwargs): # pylint: disable=useless-parent-delegation diff --git a/tests/test_link.py b/tests/test_link.py index ed9d27a37..f96961075 100644 --- a/tests/test_link.py +++ b/tests/test_link.py @@ -1,18 +1,16 @@ -import unittest - import pytest from scrapy.link import Link -class LinkTest(unittest.TestCase): +class TestLink: def _assert_same_links(self, link1, link2): - self.assertEqual(link1, link2) - self.assertEqual(hash(link1), hash(link2)) + assert link1 == link2 + assert hash(link1) == hash(link2) def _assert_different_links(self, link1, link2): - self.assertNotEqual(link1, link2) - self.assertNotEqual(hash(link1), hash(link2)) + assert link1 != link2 + assert hash(link1) != hash(link2) def test_eq_and_hash(self): l1 = Link("http://www.example.com") diff --git a/tests/test_linkextractors.py b/tests/test_linkextractors.py index e751e0a63..1bff369af 100644 --- a/tests/test_linkextractors.py +++ b/tests/test_linkextractors.py @@ -2,7 +2,6 @@ from __future__ import annotations import pickle import re -import unittest import pytest from packaging.version import Version @@ -16,175 +15,139 @@ from tests import get_testdata # a hack to skip base class tests in pytest class Base: - class LinkExtractorTestCase(unittest.TestCase): + class TestLinkExtractorBase: extractor_cls: type | None = None - def setUp(self): + def setup_method(self): body = get_testdata("link_extractor", "linkextractor.html") self.response = HtmlResponse(url="http://example.com/index", body=body) def test_urls_type(self): """Test that the resulting urls are str objects""" lx = self.extractor_cls() - self.assertTrue( - all( - isinstance(link.url, str) - for link in lx.extract_links(self.response) - ) + assert all( + isinstance(link.url, str) for link in lx.extract_links(self.response) ) def test_extract_all_links(self): lx = self.extractor_cls() page4_url = "http://example.com/page%204.html" - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - Link(url="http://example.com/sample3.html", text="sample 3 text"), - Link( - url="http://example.com/sample3.html#foo", - text="sample 3 repetition with fragment", - ), - Link(url="http://www.google.com/something", text=""), - Link(url="http://example.com/innertag.html", text="inner tag"), - Link(url=page4_url, text="href with whitespaces"), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + Link(url="http://example.com/sample3.html", text="sample 3 text"), + Link( + url="http://example.com/sample3.html#foo", + text="sample 3 repetition with fragment", + ), + Link(url="http://www.google.com/something", text=""), + Link(url="http://example.com/innertag.html", text="inner tag"), + Link(url=page4_url, text="href with whitespaces"), + ] def test_extract_filter_allow(self): lx = self.extractor_cls(allow=("sample",)) - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - Link(url="http://example.com/sample3.html", text="sample 3 text"), - Link( - url="http://example.com/sample3.html#foo", - text="sample 3 repetition with fragment", - ), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + Link(url="http://example.com/sample3.html", text="sample 3 text"), + Link( + url="http://example.com/sample3.html#foo", + text="sample 3 repetition with fragment", + ), + ] def test_extract_filter_allow_with_duplicates(self): lx = self.extractor_cls(allow=("sample",), unique=False) - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - Link(url="http://example.com/sample3.html", text="sample 3 text"), - Link( - url="http://example.com/sample3.html", - text="sample 3 repetition", - ), - Link( - url="http://example.com/sample3.html", - text="sample 3 repetition", - ), - Link( - url="http://example.com/sample3.html#foo", - text="sample 3 repetition with fragment", - ), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + Link(url="http://example.com/sample3.html", text="sample 3 text"), + Link( + url="http://example.com/sample3.html", + text="sample 3 repetition", + ), + Link( + url="http://example.com/sample3.html", + text="sample 3 repetition", + ), + Link( + url="http://example.com/sample3.html#foo", + text="sample 3 repetition with fragment", + ), + ] def test_extract_filter_allow_with_duplicates_canonicalize(self): lx = self.extractor_cls(allow=("sample",), unique=False, canonicalize=True) - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - Link(url="http://example.com/sample3.html", text="sample 3 text"), - Link( - url="http://example.com/sample3.html", - text="sample 3 repetition", - ), - Link( - url="http://example.com/sample3.html", - text="sample 3 repetition", - ), - Link( - url="http://example.com/sample3.html", - text="sample 3 repetition with fragment", - ), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + Link(url="http://example.com/sample3.html", text="sample 3 text"), + Link( + url="http://example.com/sample3.html", + text="sample 3 repetition", + ), + Link( + url="http://example.com/sample3.html", + text="sample 3 repetition", + ), + Link( + url="http://example.com/sample3.html", + text="sample 3 repetition with fragment", + ), + ] def test_extract_filter_allow_no_duplicates_canonicalize(self): lx = self.extractor_cls(allow=("sample",), unique=True, canonicalize=True) - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - Link(url="http://example.com/sample3.html", text="sample 3 text"), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + Link(url="http://example.com/sample3.html", text="sample 3 text"), + ] def test_extract_filter_allow_and_deny(self): lx = self.extractor_cls(allow=("sample",), deny=("3",)) - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + ] def test_extract_filter_allowed_domains(self): lx = self.extractor_cls(allow_domains=("google.com",)) - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://www.google.com/something", text=""), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://www.google.com/something", text=""), + ] def test_extraction_using_single_values(self): """Test the extractor's behaviour among different situations""" lx = self.extractor_cls(allow="sample") - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - Link(url="http://example.com/sample3.html", text="sample 3 text"), - Link( - url="http://example.com/sample3.html#foo", - text="sample 3 repetition with fragment", - ), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + Link(url="http://example.com/sample3.html", text="sample 3 text"), + Link( + url="http://example.com/sample3.html#foo", + text="sample 3 repetition with fragment", + ), + ] lx = self.extractor_cls(allow="sample", deny="3") - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + ] lx = self.extractor_cls(allow_domains="google.com") - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://www.google.com/something", text=""), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://www.google.com/something", text=""), + ] lx = self.extractor_cls(deny_domains="example.com") - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://www.google.com/something", text=""), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://www.google.com/something", text=""), + ] def test_nofollow(self): """Test the extractor's behaviour for links with rel='nofollow'""" @@ -210,47 +173,44 @@ class Base: response = HtmlResponse("http://example.org/somepage/index.html", body=html) lx = self.extractor_cls() - self.assertEqual( - lx.extract_links(response), - [ - Link(url="http://example.org/about.html", text="About us"), - Link(url="http://example.org/follow.html", text="Follow this link"), - Link( - url="http://example.org/nofollow.html", - text="Dont follow this one", - nofollow=True, - ), - Link( - url="http://example.org/nofollow2.html", - text="Choose to follow or not", - ), - Link( - url="http://google.com/something", - text="External link not to follow", - nofollow=True, - ), - ], - ) + assert lx.extract_links(response) == [ + Link(url="http://example.org/about.html", text="About us"), + Link(url="http://example.org/follow.html", text="Follow this link"), + Link( + url="http://example.org/nofollow.html", + text="Dont follow this one", + nofollow=True, + ), + Link( + url="http://example.org/nofollow2.html", + text="Choose to follow or not", + ), + Link( + url="http://google.com/something", + text="External link not to follow", + nofollow=True, + ), + ] def test_matches(self): url1 = "http://lotsofstuff.com/stuff1/index" url2 = "http://evenmorestuff.com/uglystuff/index" lx = self.extractor_cls(allow=(r"stuff1",)) - self.assertTrue(lx.matches(url1)) - self.assertFalse(lx.matches(url2)) + assert lx.matches(url1) + assert not lx.matches(url2) lx = self.extractor_cls(deny=(r"uglystuff",)) - self.assertTrue(lx.matches(url1)) - self.assertFalse(lx.matches(url2)) + assert lx.matches(url1) + assert not lx.matches(url2) lx = self.extractor_cls(allow_domains=("evenmorestuff.com",)) - self.assertFalse(lx.matches(url1)) - self.assertTrue(lx.matches(url2)) + assert not lx.matches(url1) + assert lx.matches(url2) lx = self.extractor_cls(deny_domains=("lotsofstuff.com",)) - self.assertFalse(lx.matches(url1)) - self.assertTrue(lx.matches(url2)) + assert not lx.matches(url1) + assert lx.matches(url2) lx = self.extractor_cls( allow=["blah1"], @@ -258,20 +218,17 @@ class Base: allow_domains=["blah1.com"], deny_domains=["blah2.com"], ) - self.assertTrue(lx.matches("http://blah1.com/blah1")) - self.assertFalse(lx.matches("http://blah1.com/blah2")) - self.assertFalse(lx.matches("http://blah2.com/blah1")) - self.assertFalse(lx.matches("http://blah2.com/blah2")) + assert lx.matches("http://blah1.com/blah1") + assert not lx.matches("http://blah1.com/blah2") + assert not lx.matches("http://blah2.com/blah1") + assert not lx.matches("http://blah2.com/blah2") def test_restrict_xpaths(self): lx = self.extractor_cls(restrict_xpaths=('//div[@id="subwrapper"]',)) - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + ] def test_restrict_xpaths_encoding(self): """Test restrict_xpaths with encodings""" @@ -291,10 +248,9 @@ class Base: ) lx = self.extractor_cls(restrict_xpaths="//div[@class='links']") - self.assertEqual( - lx.extract_links(response), - [Link(url="http://example.org/about.html", text="About us\xa3")], - ) + assert lx.extract_links(response) == [ + Link(url="http://example.org/about.html", text="About us\xa3") + ] def test_restrict_xpaths_with_html_entities(self): html = b'

text

' @@ -304,47 +260,40 @@ class Base: encoding="iso8859-15", ) links = self.extractor_cls(restrict_xpaths="//p").extract_links(response) - self.assertEqual( - links, [Link(url="http://example.org/%E2%99%A5/you?c=%A4", text="text")] - ) + assert links == [ + Link(url="http://example.org/%E2%99%A5/you?c=%A4", text="text") + ] def test_restrict_xpaths_concat_in_handle_data(self): """html entities cause SGMLParser to call handle_data hook twice""" body = b"""
>\xbe\xa9<\xb6\xab""" response = HtmlResponse("http://example.org", body=body, encoding="gb18030") lx = self.extractor_cls(restrict_xpaths="//div") - self.assertEqual( - lx.extract_links(response), - [ - Link( - url="http://example.org/foo", - text=">\u4eac<\u4e1c", - fragment="", - nofollow=False, - ) - ], - ) + assert lx.extract_links(response) == [ + Link( + url="http://example.org/foo", + text=">\u4eac<\u4e1c", + fragment="", + nofollow=False, + ) + ] def test_restrict_css(self): lx = self.extractor_cls(restrict_css=("#subwrapper a",)) - self.assertEqual( - lx.extract_links(self.response), - [Link(url="http://example.com/sample2.html", text="sample 2")], - ) + assert lx.extract_links(self.response) == [ + Link(url="http://example.com/sample2.html", text="sample 2") + ] def test_restrict_css_and_restrict_xpaths_together(self): lx = self.extractor_cls( restrict_xpaths=('//div[@id="subwrapper"]',), restrict_css=("#subwrapper + a",), ) - self.assertEqual( - list(lx.extract_links(self.response)), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - Link(url="http://example.com/sample3.html", text="sample 3 text"), - ], - ) + assert list(lx.extract_links(self.response)) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + Link(url="http://example.com/sample3.html", text="sample 3 text"), + ] def test_area_tag_with_unicode_present(self): body = b"""\xbe\xa9""" @@ -353,17 +302,14 @@ class Base: lx.extract_links(response) lx.extract_links(response) lx.extract_links(response) - self.assertEqual( - lx.extract_links(response), - [ - Link( - url="http://example.org/foo", - text="", - fragment="", - nofollow=False, - ) - ], - ) + assert lx.extract_links(response) == [ + Link( + url="http://example.org/foo", + text="", + fragment="", + nofollow=False, + ) + ] def test_encoded_url(self): body = b"""
BinB""" @@ -371,17 +317,14 @@ class Base: "http://known.fm/AC%2FDC/", body=body, encoding="utf8" ) lx = self.extractor_cls() - self.assertEqual( - lx.extract_links(response), - [ - Link( - url="http://known.fm/AC%2FDC/?page=2", - text="BinB", - fragment="", - nofollow=False, - ), - ], - ) + assert lx.extract_links(response) == [ + Link( + url="http://known.fm/AC%2FDC/?page=2", + text="BinB", + fragment="", + nofollow=False, + ), + ] def test_encoded_url_in_restricted_xpath(self): body = b"""
BinB""" @@ -389,38 +332,29 @@ class Base: "http://known.fm/AC%2FDC/", body=body, encoding="utf8" ) lx = self.extractor_cls(restrict_xpaths="//div") - self.assertEqual( - lx.extract_links(response), - [ - Link( - url="http://known.fm/AC%2FDC/?page=2", - text="BinB", - fragment="", - nofollow=False, - ), - ], - ) + assert lx.extract_links(response) == [ + Link( + url="http://known.fm/AC%2FDC/?page=2", + text="BinB", + fragment="", + nofollow=False, + ), + ] def test_ignored_extensions(self): # jpg is ignored by default html = b"""asd and """ response = HtmlResponse("http://example.org/", body=html) lx = self.extractor_cls() - self.assertEqual( - lx.extract_links(response), - [ - Link(url="http://example.org/page.html", text="asd"), - ], - ) + assert lx.extract_links(response) == [ + Link(url="http://example.org/page.html", text="asd"), + ] # override denied extensions lx = self.extractor_cls(deny_extensions=["html"]) - self.assertEqual( - lx.extract_links(response), - [ - Link(url="http://example.org/photo.jpg"), - ], - ) + assert lx.extract_links(response) == [ + Link(url="http://example.org/photo.jpg"), + ] def test_process_value(self): """Test restrict_xpaths with encodings""" @@ -439,10 +373,9 @@ class Base: return m.group(1) if m else None lx = self.extractor_cls(process_value=process_value) - self.assertEqual( - lx.extract_links(response), - [Link(url="http://example.org/other/page.html", text="Text")], - ) + assert lx.extract_links(response) == [ + Link(url="http://example.org/other/page.html", text="Text") + ] def test_base_url_with_restrict_xpaths(self): html = b"""Page title<title><base href="http://otherdomain.com/base/" /> @@ -450,53 +383,46 @@ class Base: </body></html>""" response = HtmlResponse("http://example.org/somepage/index.html", body=html) lx = self.extractor_cls(restrict_xpaths="//p") - self.assertEqual( - lx.extract_links(response), - [Link(url="http://otherdomain.com/base/item/12.html", text="Item 12")], - ) + assert lx.extract_links(response) == [ + Link(url="http://otherdomain.com/base/item/12.html", text="Item 12") + ] def test_attrs(self): lx = self.extractor_cls(attrs="href") page4_url = "http://example.com/page%204.html" - self.assertEqual( - lx.extract_links(self.response), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - Link(url="http://example.com/sample3.html", text="sample 3 text"), - Link( - url="http://example.com/sample3.html#foo", - text="sample 3 repetition with fragment", - ), - Link(url="http://www.google.com/something", text=""), - Link(url="http://example.com/innertag.html", text="inner tag"), - Link(url=page4_url, text="href with whitespaces"), - ], - ) + assert lx.extract_links(self.response) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + Link(url="http://example.com/sample3.html", text="sample 3 text"), + Link( + url="http://example.com/sample3.html#foo", + text="sample 3 repetition with fragment", + ), + Link(url="http://www.google.com/something", text=""), + Link(url="http://example.com/innertag.html", text="inner tag"), + Link(url=page4_url, text="href with whitespaces"), + ] lx = self.extractor_cls( attrs=("href", "src"), tags=("a", "area", "img"), deny_extensions=() ) - self.assertEqual( - lx.extract_links(self.response), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - Link(url="http://example.com/sample2.jpg", text=""), - Link(url="http://example.com/sample3.html", text="sample 3 text"), - Link( - url="http://example.com/sample3.html#foo", - text="sample 3 repetition with fragment", - ), - Link(url="http://www.google.com/something", text=""), - Link(url="http://example.com/innertag.html", text="inner tag"), - Link(url=page4_url, text="href with whitespaces"), - ], - ) + assert lx.extract_links(self.response) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + Link(url="http://example.com/sample2.jpg", text=""), + Link(url="http://example.com/sample3.html", text="sample 3 text"), + Link( + url="http://example.com/sample3.html#foo", + text="sample 3 repetition with fragment", + ), + Link(url="http://www.google.com/something", text=""), + Link(url="http://example.com/innertag.html", text="inner tag"), + Link(url=page4_url, text="href with whitespaces"), + ] lx = self.extractor_cls(attrs=None) - self.assertEqual(lx.extract_links(self.response), []) + assert lx.extract_links(self.response) == [] def test_tags(self): html = ( @@ -506,43 +432,31 @@ class Base: response = HtmlResponse("http://example.com/index.html", body=html) lx = self.extractor_cls(tags=None) - self.assertEqual(lx.extract_links(response), []) + assert lx.extract_links(response) == [] lx = self.extractor_cls() - self.assertEqual( - lx.extract_links(response), - [ - Link(url="http://example.com/sample1.html", text=""), - Link(url="http://example.com/sample2.html", text="sample 2"), - ], - ) + assert lx.extract_links(response) == [ + Link(url="http://example.com/sample1.html", text=""), + Link(url="http://example.com/sample2.html", text="sample 2"), + ] lx = self.extractor_cls(tags="area") - self.assertEqual( - lx.extract_links(response), - [ - Link(url="http://example.com/sample1.html", text=""), - ], - ) + assert lx.extract_links(response) == [ + Link(url="http://example.com/sample1.html", text=""), + ] lx = self.extractor_cls(tags="a") - self.assertEqual( - lx.extract_links(response), - [ - Link(url="http://example.com/sample2.html", text="sample 2"), - ], - ) + assert lx.extract_links(response) == [ + Link(url="http://example.com/sample2.html", text="sample 2"), + ] lx = self.extractor_cls( tags=("a", "img"), attrs=("href", "src"), deny_extensions=() ) - self.assertEqual( - lx.extract_links(response), - [ - Link(url="http://example.com/sample2.html", text="sample 2"), - Link(url="http://example.com/sample2.jpg", text=""), - ], - ) + assert lx.extract_links(response) == [ + Link(url="http://example.com/sample2.html", text="sample 2"), + Link(url="http://example.com/sample2.jpg", text=""), + ] def test_tags_attrs(self): html = b""" @@ -554,42 +468,36 @@ class Base: response = HtmlResponse("http://example.com/index.html", body=html) lx = self.extractor_cls(tags="div", attrs="data-url") - self.assertEqual( - lx.extract_links(response), - [ - Link( - url="http://example.com/get?id=1", - text="Item 1", - fragment="", - nofollow=False, - ), - Link( - url="http://example.com/get?id=2", - text="Item 2", - fragment="", - nofollow=False, - ), - ], - ) + assert lx.extract_links(response) == [ + Link( + url="http://example.com/get?id=1", + text="Item 1", + fragment="", + nofollow=False, + ), + Link( + url="http://example.com/get?id=2", + text="Item 2", + fragment="", + nofollow=False, + ), + ] lx = self.extractor_cls(tags=("div",), attrs=("data-url",)) - self.assertEqual( - lx.extract_links(response), - [ - Link( - url="http://example.com/get?id=1", - text="Item 1", - fragment="", - nofollow=False, - ), - Link( - url="http://example.com/get?id=2", - text="Item 2", - fragment="", - nofollow=False, - ), - ], - ) + assert lx.extract_links(response) == [ + Link( + url="http://example.com/get?id=1", + text="Item 1", + fragment="", + nofollow=False, + ), + Link( + url="http://example.com/get?id=2", + text="Item 2", + fragment="", + nofollow=False, + ), + ] def test_xhtml(self): xhtml = b""" @@ -623,78 +531,72 @@ class Base: response = HtmlResponse("http://example.com/index.xhtml", body=xhtml) lx = self.extractor_cls() - self.assertEqual( - lx.extract_links(response), - [ - Link( - url="http://example.com/about.html", - text="About us", - fragment="", - nofollow=False, - ), - Link( - url="http://example.com/follow.html", - text="Follow this link", - fragment="", - nofollow=False, - ), - Link( - url="http://example.com/nofollow.html", - text="Dont follow this one", - fragment="", - nofollow=True, - ), - Link( - url="http://example.com/nofollow2.html", - text="Choose to follow or not", - fragment="", - nofollow=False, - ), - Link( - url="http://google.com/something", - text="External link not to follow", - nofollow=True, - ), - ], - ) + assert lx.extract_links(response) == [ + Link( + url="http://example.com/about.html", + text="About us", + fragment="", + nofollow=False, + ), + Link( + url="http://example.com/follow.html", + text="Follow this link", + fragment="", + nofollow=False, + ), + Link( + url="http://example.com/nofollow.html", + text="Dont follow this one", + fragment="", + nofollow=True, + ), + Link( + url="http://example.com/nofollow2.html", + text="Choose to follow or not", + fragment="", + nofollow=False, + ), + Link( + url="http://google.com/something", + text="External link not to follow", + nofollow=True, + ), + ] response = XmlResponse("http://example.com/index.xhtml", body=xhtml) lx = self.extractor_cls() - self.assertEqual( - lx.extract_links(response), - [ - Link( - url="http://example.com/about.html", - text="About us", - fragment="", - nofollow=False, - ), - Link( - url="http://example.com/follow.html", - text="Follow this link", - fragment="", - nofollow=False, - ), - Link( - url="http://example.com/nofollow.html", - text="Dont follow this one", - fragment="", - nofollow=True, - ), - Link( - url="http://example.com/nofollow2.html", - text="Choose to follow or not", - fragment="", - nofollow=False, - ), - Link( - url="http://google.com/something", - text="External link not to follow", - nofollow=True, - ), - ], - ) + assert lx.extract_links(response) == [ + Link( + url="http://example.com/about.html", + text="About us", + fragment="", + nofollow=False, + ), + Link( + url="http://example.com/follow.html", + text="Follow this link", + fragment="", + nofollow=False, + ), + Link( + url="http://example.com/nofollow.html", + text="Dont follow this one", + fragment="", + nofollow=True, + ), + Link( + url="http://example.com/nofollow2.html", + text="Choose to follow or not", + fragment="", + nofollow=False, + ), + Link( + url="http://google.com/something", + text="External link not to follow", + nofollow=True, + ), + ] def test_link_wrong_href(self): html = b""" @@ -704,21 +606,18 @@ class Base: """ response = HtmlResponse("http://example.org/index.html", body=html) lx = self.extractor_cls() - self.assertEqual( - list(lx.extract_links(response)), - [ - Link( - url="http://example.org/item1.html", - text="Item 1", - nofollow=False, - ), - Link( - url="http://example.org/item3.html", - text="Item 3", - nofollow=False, - ), - ], - ) + assert list(lx.extract_links(response)) == [ + Link( + url="http://example.org/item1.html", + text="Item 1", + nofollow=False, + ), + Link( + url="http://example.org/item3.html", + text="Item 3", + nofollow=False, + ), + ] def test_ftp_links(self): body = b""" @@ -729,21 +628,18 @@ class Base: "http://www.example.com/index.html", body=body, encoding="utf8" ) lx = self.extractor_cls() - self.assertEqual( - lx.extract_links(response), - [ - Link( - url="ftp://www.external.com/", - text="An Item", - fragment="", - nofollow=False, - ), - ], - ) + assert lx.extract_links(response) == [ + Link( + url="ftp://www.external.com/", + text="An Item", + fragment="", + nofollow=False, + ), + ] def test_pickle_extractor(self): lx = self.extractor_cls() - self.assertIsInstance(pickle.loads(pickle.dumps(lx)), self.extractor_cls) + assert isinstance(pickle.loads(pickle.dumps(lx)), self.extractor_cls) def test_link_extractor_aggregation(self): """When a parameter like restrict_css is used, the underlying @@ -770,14 +666,11 @@ class Base: """, ) actual = lx.extract_links(response) - self.assertEqual( - actual, - [ - Link(url="https://example.com/a", text="a1"), - Link(url="https://example.com/b?a=1&b=2", text="b1"), - Link(url="https://example.com/b?b=2&a=1", text="b2"), - ], - ) + assert actual == [ + Link(url="https://example.com/a", text="a1"), + Link(url="https://example.com/b?a=1&b=2", text="b1"), + Link(url="https://example.com/b?b=2&a=1", text="b2"), + ] # unique=True (default), canonicalize=True lx = self.extractor_cls(restrict_css=("div",), canonicalize=True) @@ -795,13 +688,10 @@ class Base: """, ) actual = lx.extract_links(response) - self.assertEqual( - actual, - [ - Link(url="https://example.com/a", text="a1"), - Link(url="https://example.com/b?a=1&b=2", text="b1"), - ], - ) + assert actual == [ + Link(url="https://example.com/a", text="a1"), + Link(url="https://example.com/b?a=1&b=2", text="b1"), + ] # unique=False, canonicalize=False (default) lx = self.extractor_cls(restrict_css=("div",), unique=False) @@ -819,15 +709,12 @@ class Base: """, ) actual = lx.extract_links(response) - self.assertEqual( - actual, - [ - Link(url="https://example.com/a", text="a1"), - Link(url="https://example.com/b?a=1&b=2", text="b1"), - Link(url="https://example.com/a", text="a2"), - Link(url="https://example.com/b?b=2&a=1", text="b2"), - ], - ) + assert actual == [ + Link(url="https://example.com/a", text="a1"), + Link(url="https://example.com/b?a=1&b=2", text="b1"), + Link(url="https://example.com/a", text="a2"), + Link(url="https://example.com/b?b=2&a=1", text="b2"), + ] # unique=False, canonicalize=True lx = self.extractor_cls( @@ -847,18 +734,15 @@ class Base: """, ) actual = lx.extract_links(response) - self.assertEqual( - actual, - [ - Link(url="https://example.com/a", text="a1"), - Link(url="https://example.com/b?a=1&b=2", text="b1"), - Link(url="https://example.com/a", text="a2"), - Link(url="https://example.com/b?a=1&b=2", text="b2"), - ], - ) + assert actual == [ + Link(url="https://example.com/a", text="a1"), + Link(url="https://example.com/b?a=1&b=2", text="b1"), + Link(url="https://example.com/a", text="a2"), + Link(url="https://example.com/b?a=1&b=2", text="b2"), + ] -class LxmlLinkExtractorTestCase(Base.LinkExtractorTestCase): +class TestLxmlLinkExtractor(Base.TestLinkExtractorBase): extractor_cls = LxmlLinkExtractor def test_link_wrong_href(self): @@ -869,17 +753,10 @@ class LxmlLinkExtractorTestCase(Base.LinkExtractorTestCase): """ response = HtmlResponse("http://example.org/index.html", body=html) lx = self.extractor_cls() - self.assertEqual( - list(lx.extract_links(response)), - [ - Link( - url="http://example.org/item1.html", text="Item 1", nofollow=False - ), - Link( - url="http://example.org/item3.html", text="Item 3", nofollow=False - ), - ], - ) + assert list(lx.extract_links(response)) == [ + Link(url="http://example.org/item1.html", text="Item 1", nofollow=False), + Link(url="http://example.org/item3.html", text="Item 3", nofollow=False), + ] def test_link_restrict_text(self): html = b""" @@ -890,45 +767,36 @@ class LxmlLinkExtractorTestCase(Base.LinkExtractorTestCase): response = HtmlResponse("http://example.org/index.html", body=html) # Simple text inclusion test lx = self.extractor_cls(restrict_text="dog") - self.assertEqual( - list(lx.extract_links(response)), - [ - Link( - url="http://example.org/item2.html", - text="Pic of a dog", - nofollow=False, - ), - ], - ) + assert list(lx.extract_links(response)) == [ + Link( + url="http://example.org/item2.html", + text="Pic of a dog", + nofollow=False, + ), + ] # Unique regex test lx = self.extractor_cls(restrict_text=r"of.*dog") - self.assertEqual( - list(lx.extract_links(response)), - [ - Link( - url="http://example.org/item2.html", - text="Pic of a dog", - nofollow=False, - ), - ], - ) + assert list(lx.extract_links(response)) == [ + Link( + url="http://example.org/item2.html", + text="Pic of a dog", + nofollow=False, + ), + ] # Multiple regex test lx = self.extractor_cls(restrict_text=[r"of.*dog", r"of.*cat"]) - self.assertEqual( - list(lx.extract_links(response)), - [ - Link( - url="http://example.org/item1.html", - text="Pic of a cat", - nofollow=False, - ), - Link( - url="http://example.org/item2.html", - text="Pic of a dog", - nofollow=False, - ), - ], - ) + assert list(lx.extract_links(response)) == [ + Link( + url="http://example.org/item1.html", + text="Pic of a cat", + nofollow=False, + ), + Link( + url="http://example.org/item2.html", + text="Pic of a dog", + nofollow=False, + ), + ] @pytest.mark.skipif( Version(w3lib_version) < Version("2.0.0"), @@ -945,30 +813,27 @@ class LxmlLinkExtractorTestCase(Base.LinkExtractorTestCase): """ response = HtmlResponse("http://example.org/index.html", body=html) lx = self.extractor_cls() - self.assertEqual( - list(lx.extract_links(response)), - [ - Link( - url="http://example.org/item2.html", - text="Good Link", - nofollow=False, - ), - Link( - url="http://example.org/item3.html", - text="Good Link 2", - nofollow=False, - ), - ], - ) + assert list(lx.extract_links(response)) == [ + Link( + url="http://example.org/item2.html", + text="Good Link", + nofollow=False, + ), + Link( + url="http://example.org/item3.html", + text="Good Link 2", + nofollow=False, + ), + ] def test_link_allowed_is_false_with_empty_url(self): bad_link = Link("") - self.assertFalse(LxmlLinkExtractor()._link_allowed(bad_link)) + assert not LxmlLinkExtractor()._link_allowed(bad_link) def test_link_allowed_is_false_with_bad_url_prefix(self): bad_link = Link("htp://should_be_http.example") - self.assertFalse(LxmlLinkExtractor()._link_allowed(bad_link)) + assert not LxmlLinkExtractor()._link_allowed(bad_link) def test_link_allowed_is_false_with_missing_url_prefix(self): bad_link = Link("should_have_prefix.example") - self.assertFalse(LxmlLinkExtractor()._link_allowed(bad_link)) + assert not LxmlLinkExtractor()._link_allowed(bad_link) diff --git a/tests/test_pipeline_crawl.py b/tests/test_pipeline_crawl.py index 84d714e5c..162dfdaf4 100644 --- a/tests/test_pipeline_crawl.py +++ b/tests/test_pipeline_crawl.py @@ -53,7 +53,7 @@ class RedirectedMediaDownloadSpider(MediaDownloadSpider): ) -class FileDownloadCrawlTestCase(TestCase): +class TestFileDownloadCrawl(TestCase): pipeline_class = "scrapy.pipelines.files.FilesPipeline" store_setting_key = "FILES_STORE" media_key = "files" @@ -98,52 +98,46 @@ class FileDownloadCrawlTestCase(TestCase): return crawler def _assert_files_downloaded(self, items, logs): - self.assertEqual(len(items), 1) - self.assertIn(self.media_key, items[0]) + assert len(items) == 1 + assert self.media_key in items[0] # check that logs show the expected number of successful file downloads file_dl_success = "File (downloaded): Downloaded file from" - self.assertEqual(logs.count(file_dl_success), 3) + assert logs.count(file_dl_success) == 3 # check that the images/files status is `downloaded` for item in items: for i in item[self.media_key]: - self.assertEqual(i["status"], "downloaded") + assert i["status"] == "downloaded" # check that the images/files checksums are what we know they should be if self.expected_checksums is not None: checksums = {i["checksum"] for item in items for i in item[self.media_key]} - self.assertEqual(checksums, self.expected_checksums) + assert checksums == self.expected_checksums # check that the image files where actually written to the media store for item in items: for i in item[self.media_key]: - self.assertTrue((self.tmpmediastore / i["path"]).exists()) + assert (self.tmpmediastore / i["path"]).exists() def _assert_files_download_failure(self, crawler, items, code, logs): # check that the item does NOT have the "images/files" field populated - self.assertEqual(len(items), 1) - self.assertIn(self.media_key, items[0]) - self.assertFalse(items[0][self.media_key]) + assert len(items) == 1 + assert self.media_key in items[0] + assert not items[0][self.media_key] # check that there was 1 successful fetch and 3 other responses with non-200 code - self.assertEqual( - crawler.stats.get_value("downloader/request_method_count/GET"), 4 - ) - self.assertEqual(crawler.stats.get_value("downloader/response_count"), 4) - self.assertEqual( - crawler.stats.get_value("downloader/response_status_count/200"), 1 - ) - self.assertEqual( - crawler.stats.get_value(f"downloader/response_status_count/{code}"), 3 - ) + assert crawler.stats.get_value("downloader/request_method_count/GET") == 4 + assert crawler.stats.get_value("downloader/response_count") == 4 + assert crawler.stats.get_value("downloader/response_status_count/200") == 1 + assert crawler.stats.get_value(f"downloader/response_status_count/{code}") == 3 # check that logs do show the failure on the file downloads file_dl_failure = f"File (code: {code}): Error downloading file from" - self.assertEqual(logs.count(file_dl_failure), 3) + assert logs.count(file_dl_failure) == 3 # check that no files were written to the media store - self.assertEqual(list(self.tmpmediastore.iterdir()), []) + assert not list(self.tmpmediastore.iterdir()) @defer.inlineCallbacks def test_download_media(self): @@ -193,9 +187,7 @@ class FileDownloadCrawlTestCase(TestCase): mockserver=self.mockserver, ) self._assert_files_downloaded(self.items, str(log)) - self.assertEqual( - crawler.stats.get_value("downloader/response_status_count/302"), 3 - ) + assert crawler.stats.get_value("downloader/response_status_count/302") == 3 @defer.inlineCallbacks def test_download_media_file_path_error(self): @@ -218,7 +210,7 @@ class FileDownloadCrawlTestCase(TestCase): media_urls_key=self.media_urls_key, mockserver=self.mockserver, ) - self.assertIn("ZeroDivisionError", str(log)) + assert "ZeroDivisionError" in str(log) skip_pillow: str | None @@ -230,7 +222,7 @@ else: skip_pillow = None -class ImageDownloadCrawlTestCase(FileDownloadCrawlTestCase): +class ImageDownloadCrawlTestCase(TestFileDownloadCrawl): skip = skip_pillow pipeline_class = "scrapy.pipelines.images.ImagesPipeline" diff --git a/tests/test_pipeline_files.py b/tests/test_pipeline_files.py index 05fd17207..e515c16a0 100644 --- a/tests/test_pipeline_files.py +++ b/tests/test_pipeline_files.py @@ -77,7 +77,7 @@ def get_ftp_content_and_delete( return b"".join(ftp_data) -class FilesPipelineTestCase(unittest.TestCase): +class TestFilesPipeline(unittest.TestCase): def setUp(self): self.tempdir = mkdtemp() settings_dict = {"FILES_STORE": self.tempdir} @@ -91,73 +91,73 @@ class FilesPipelineTestCase(unittest.TestCase): def test_file_path(self): file_path = self.pipeline.file_path - self.assertEqual( - file_path(Request("https://dev.mydeco.com/mydeco.pdf")), - "full/c9b564df929f4bc635bdd19fde4f3d4847c757c5.pdf", + assert ( + file_path(Request("https://dev.mydeco.com/mydeco.pdf")) + == "full/c9b564df929f4bc635bdd19fde4f3d4847c757c5.pdf" ) - self.assertEqual( + assert ( file_path( Request( "http://www.maddiebrown.co.uk///catalogue-items//image_54642_12175_95307.txt" ) - ), - "full/4ce274dd83db0368bafd7e406f382ae088e39219.txt", + ) + == "full/4ce274dd83db0368bafd7e406f382ae088e39219.txt" ) - self.assertEqual( + assert ( file_path( Request("https://dev.mydeco.com/two/dirs/with%20spaces%2Bsigns.doc") - ), - "full/94ccc495a17b9ac5d40e3eabf3afcb8c2c9b9e1a.doc", + ) + == "full/94ccc495a17b9ac5d40e3eabf3afcb8c2c9b9e1a.doc" ) - self.assertEqual( + assert ( file_path( Request( "http://www.dfsonline.co.uk/get_prod_image.php?img=status_0907_mdm.jpg" ) - ), - "full/4507be485f38b0da8a0be9eb2e1dfab8a19223f2.jpg", + ) + == "full/4507be485f38b0da8a0be9eb2e1dfab8a19223f2.jpg" ) - self.assertEqual( - file_path(Request("http://www.dorma.co.uk/images/product_details/2532/")), - "full/97ee6f8a46cbbb418ea91502fd24176865cf39b2", + assert ( + file_path(Request("http://www.dorma.co.uk/images/product_details/2532/")) + == "full/97ee6f8a46cbbb418ea91502fd24176865cf39b2" ) - self.assertEqual( - file_path(Request("http://www.dorma.co.uk/images/product_details/2532")), - "full/244e0dd7d96a3b7b01f54eded250c9e272577aa1", + assert ( + file_path(Request("http://www.dorma.co.uk/images/product_details/2532")) + == "full/244e0dd7d96a3b7b01f54eded250c9e272577aa1" ) - self.assertEqual( + assert ( file_path( Request("http://www.dorma.co.uk/images/product_details/2532"), response=Response("http://www.dorma.co.uk/images/product_details/2532"), info=object(), - ), - "full/244e0dd7d96a3b7b01f54eded250c9e272577aa1", + ) + == "full/244e0dd7d96a3b7b01f54eded250c9e272577aa1" ) - self.assertEqual( + assert ( file_path( Request( "http://www.dfsonline.co.uk/get_prod_image.php?img=status_0907_mdm.jpg.bohaha" ) - ), - "full/76c00cef2ef669ae65052661f68d451162829507", + ) + == "full/76c00cef2ef669ae65052661f68d451162829507" ) - self.assertEqual( + assert ( file_path( Request( "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAR0AAACxCAMAAADOHZloAAACClBMVEX/\ //+F0tzCwMK76ZKQ21AMqr7oAAC96JvD5aWM2kvZ78J0N7fmAAC46Y4Ap7y" ) - ), - "full/178059cbeba2e34120a67f2dc1afc3ecc09b61cb.png", + ) + == "full/178059cbeba2e34120a67f2dc1afc3ecc09b61cb.png" ) def test_fs_store(self): assert isinstance(self.pipeline.store, FSFilesStore) - self.assertEqual(self.pipeline.store.basedir, self.tempdir) + assert self.pipeline.store.basedir == self.tempdir path = "some/image/key.jpg" fullpath = Path(self.tempdir, "some", "image", "key.jpg") - self.assertEqual(self.pipeline.store._get_filesystem_path(path), fullpath) + assert self.pipeline.store._get_filesystem_path(path) == fullpath @defer.inlineCallbacks def test_file_not_expired(self): @@ -180,8 +180,8 @@ class FilesPipelineTestCase(unittest.TestCase): p.start() result = yield self.pipeline.process_item(item, None) - self.assertEqual(result["files"][0]["checksum"], "abc") - self.assertEqual(result["files"][0]["status"], "uptodate") + assert result["files"][0]["checksum"] == "abc" + assert result["files"][0]["status"] == "uptodate" for p in patchers: p.stop() @@ -211,8 +211,8 @@ class FilesPipelineTestCase(unittest.TestCase): p.start() result = yield self.pipeline.process_item(item, None) - self.assertNotEqual(result["files"][0]["checksum"], "abc") - self.assertEqual(result["files"][0]["status"], "downloaded") + assert result["files"][0]["checksum"] != "abc" + assert result["files"][0]["status"] == "downloaded" for p in patchers: p.stop() @@ -242,8 +242,8 @@ class FilesPipelineTestCase(unittest.TestCase): p.start() result = yield self.pipeline.process_item(item, None) - self.assertNotEqual(result["files"][0]["checksum"], "abc") - self.assertEqual(result["files"][0]["status"], "cached") + assert result["files"][0]["checksum"] != "abc" + assert result["files"][0]["status"] == "cached" for p in patchers: p.stop() @@ -262,14 +262,14 @@ class FilesPipelineTestCase(unittest.TestCase): ).file_path item = {"path": "path-to-store-file"} request = Request("http://example.com") - self.assertEqual(file_path(request, item=item), "full/path-to-store-file") + assert file_path(request, item=item) == "full/path-to-store-file" class FilesPipelineTestCaseFieldsMixin: - def setUp(self): + def setup_method(self): self.tempdir = mkdtemp() - def tearDown(self): + def teardown_method(self): rmtree(self.tempdir) def test_item_fields_default(self): @@ -279,12 +279,12 @@ class FilesPipelineTestCaseFieldsMixin: get_crawler(None, {"FILES_STORE": self.tempdir}) ) requests = list(pipeline.get_media_requests(item, None)) - self.assertEqual(requests[0].url, url) + assert requests[0].url == url results = [(True, {"url": url})] item = pipeline.item_completed(results, item, None) files = ItemAdapter(item).get("files") - self.assertEqual(files, [results[0][1]]) - self.assertIsInstance(item, self.item_class) + assert files == [results[0][1]] + assert isinstance(item, self.item_class) def test_item_fields_override_settings(self): url = "http://www.example.com/files/1.txt" @@ -300,17 +300,15 @@ class FilesPipelineTestCaseFieldsMixin: ) ) requests = list(pipeline.get_media_requests(item, None)) - self.assertEqual(requests[0].url, url) + assert requests[0].url == url results = [(True, {"url": url})] item = pipeline.item_completed(results, item, None) custom_files = ItemAdapter(item).get("custom_files") - self.assertEqual(custom_files, [results[0][1]]) - self.assertIsInstance(item, self.item_class) + assert custom_files == [results[0][1]] + assert isinstance(item, self.item_class) -class FilesPipelineTestCaseFieldsDict( - FilesPipelineTestCaseFieldsMixin, unittest.TestCase -): +class TestFilesPipelineFieldsDict(FilesPipelineTestCaseFieldsMixin): item_class = dict @@ -324,9 +322,7 @@ class FilesPipelineTestItem(Item): custom_files = Field() -class FilesPipelineTestCaseFieldsItem( - FilesPipelineTestCaseFieldsMixin, unittest.TestCase -): +class TestFilesPipelineFieldsItem(FilesPipelineTestCaseFieldsMixin): item_class = FilesPipelineTestItem @@ -341,9 +337,7 @@ class FilesPipelineTestDataClass: custom_files: list = dataclasses.field(default_factory=list) -class FilesPipelineTestCaseFieldsDataClass( - FilesPipelineTestCaseFieldsMixin, unittest.TestCase -): +class TestFilesPipelineFieldsDataClass(FilesPipelineTestCaseFieldsMixin): item_class = FilesPipelineTestDataClass @@ -358,13 +352,11 @@ class FilesPipelineTestAttrsItem: custom_files: list[dict[str, str]] = attr.ib(default=list) -class FilesPipelineTestCaseFieldsAttrsItem( - FilesPipelineTestCaseFieldsMixin, unittest.TestCase -): +class TestFilesPipelineFieldsAttrsItem(FilesPipelineTestCaseFieldsMixin): item_class = FilesPipelineTestAttrsItem -class FilesPipelineTestCaseCustomSettings(unittest.TestCase): +class TestFilesPipelineCustomSettings: default_cls_settings = { "EXPIRES": 90, "FILES_URLS_FIELD": "file_urls", @@ -376,10 +368,10 @@ class FilesPipelineTestCaseCustomSettings(unittest.TestCase): ("FILES_RESULT_FIELD", "FILES_RESULT_FIELD", "files_result_field"), } - def setUp(self): + def setup_method(self): self.tempdir = mkdtemp() - def tearDown(self): + def teardown_method(self): rmtree(self.tempdir) def _generate_fake_settings(self, prefix=None): @@ -420,10 +412,10 @@ class FilesPipelineTestCaseCustomSettings(unittest.TestCase): one_pipeline = FilesPipeline(self.tempdir, crawler=get_crawler(None)) for pipe_attr, settings_attr, pipe_ins_attr in self.file_cls_attr_settings_map: default_value = self.default_cls_settings[pipe_attr] - self.assertEqual(getattr(one_pipeline, pipe_attr), default_value) + assert getattr(one_pipeline, pipe_attr) == default_value custom_value = custom_settings[settings_attr] - self.assertNotEqual(default_value, custom_value) - self.assertEqual(getattr(another_pipeline, pipe_ins_attr), custom_value) + assert default_value != custom_value + assert getattr(another_pipeline, pipe_ins_attr) == custom_value def test_subclass_attributes_preserved_if_no_settings(self): """ @@ -433,8 +425,8 @@ class FilesPipelineTestCaseCustomSettings(unittest.TestCase): pipe = pipe_cls.from_crawler(get_crawler(None, {"FILES_STORE": self.tempdir})) for pipe_attr, settings_attr, pipe_ins_attr in self.file_cls_attr_settings_map: custom_value = getattr(pipe, pipe_ins_attr) - self.assertNotEqual(custom_value, self.default_cls_settings[pipe_attr]) - self.assertEqual(getattr(pipe, pipe_ins_attr), getattr(pipe, pipe_attr)) + assert custom_value != self.default_cls_settings[pipe_attr] + assert getattr(pipe, pipe_ins_attr) == getattr(pipe, pipe_attr) def test_subclass_attrs_preserved_custom_settings(self): """ @@ -447,8 +439,8 @@ class FilesPipelineTestCaseCustomSettings(unittest.TestCase): for pipe_attr, settings_attr, pipe_ins_attr in self.file_cls_attr_settings_map: value = getattr(pipeline, pipe_ins_attr) setting_value = settings.get(settings_attr) - self.assertNotEqual(value, self.default_cls_settings[pipe_attr]) - self.assertEqual(value, setting_value) + assert value != self.default_cls_settings[pipe_attr] + assert value == setting_value def test_no_custom_settings_for_subclasses(self): """ @@ -465,7 +457,7 @@ class FilesPipelineTestCaseCustomSettings(unittest.TestCase): for pipe_attr, settings_attr, pipe_ins_attr in self.file_cls_attr_settings_map: # Values from settings for custom pipeline should be set on pipeline instance. custom_value = self.default_cls_settings.get(pipe_attr.upper()) - self.assertEqual(getattr(user_pipeline, pipe_ins_attr), custom_value) + assert getattr(user_pipeline, pipe_ins_attr) == custom_value def test_custom_settings_for_subclasses(self): """ @@ -484,8 +476,8 @@ class FilesPipelineTestCaseCustomSettings(unittest.TestCase): for pipe_attr, settings_attr, pipe_inst_attr in self.file_cls_attr_settings_map: # Values from settings for custom pipeline should be set on pipeline instance. custom_value = settings.get(prefix + "_" + settings_attr) - self.assertNotEqual(custom_value, self.default_cls_settings[pipe_attr]) - self.assertEqual(getattr(user_pipeline, pipe_inst_attr), custom_value) + assert custom_value != self.default_cls_settings[pipe_attr] + assert getattr(user_pipeline, pipe_inst_attr) == custom_value def test_custom_settings_and_class_attrs_for_subclasses(self): """ @@ -502,8 +494,8 @@ class FilesPipelineTestCaseCustomSettings(unittest.TestCase): pipe_inst_attr, ) in self.file_cls_attr_settings_map: custom_value = settings.get(prefix + "_" + settings_attr) - self.assertNotEqual(custom_value, self.default_cls_settings[pipe_cls_attr]) - self.assertEqual(getattr(user_pipeline, pipe_inst_attr), custom_value) + assert custom_value != self.default_cls_settings[pipe_cls_attr] + assert getattr(user_pipeline, pipe_inst_attr) == custom_value def test_cls_attrs_with_DEFAULT_prefix(self): class UserDefinedFilesPipeline(FilesPipeline): @@ -513,12 +505,13 @@ class FilesPipelineTestCaseCustomSettings(unittest.TestCase): pipeline = UserDefinedFilesPipeline.from_crawler( get_crawler(None, {"FILES_STORE": self.tempdir}) ) - self.assertEqual( - pipeline.files_result_field, - UserDefinedFilesPipeline.DEFAULT_FILES_RESULT_FIELD, + assert ( + pipeline.files_result_field + == UserDefinedFilesPipeline.DEFAULT_FILES_RESULT_FIELD ) - self.assertEqual( - pipeline.files_urls_field, UserDefinedFilesPipeline.DEFAULT_FILES_URLS_FIELD + assert ( + pipeline.files_urls_field + == UserDefinedFilesPipeline.DEFAULT_FILES_URLS_FIELD ) def test_user_defined_subclass_default_key_names(self): @@ -535,7 +528,7 @@ class FilesPipelineTestCaseCustomSettings(unittest.TestCase): for pipe_attr, settings_attr, pipe_inst_attr in self.file_cls_attr_settings_map: expected_value = settings.get(settings_attr) - self.assertEqual(getattr(pipeline_cls, pipe_inst_attr), expected_value) + assert getattr(pipeline_cls, pipe_inst_attr) == expected_value def test_file_pipeline_using_pathlike_objects(self): class CustomFilesPipelineWithPathLikeDir(FilesPipeline): @@ -546,12 +539,12 @@ class FilesPipelineTestCaseCustomSettings(unittest.TestCase): get_crawler(None, {"FILES_STORE": Path("./Temp")}) ) request = Request("http://example.com/image01.jpg") - self.assertEqual(pipeline.file_path(request), Path("subdir/image01.jpg")) + assert pipeline.file_path(request) == Path("subdir/image01.jpg") def test_files_store_constructor_with_pathlike_object(self): path = Path("./FileDir") fs_store = FSFilesStore(path) - self.assertEqual(fs_store.basedir, str(path)) + assert fs_store.basedir == str(path) @pytest.mark.requires_botocore @@ -593,13 +586,8 @@ class TestS3FilesStore(unittest.TestCase): ) stub.assert_no_pending_responses() - self.assertEqual( - buffer.method_calls, - [ - mock.call.seek(0), - # The call to read does not happen with Stubber - ], - ) + # The call to read does not happen with Stubber + assert buffer.method_calls == [mock.call.seek(0)] @defer.inlineCallbacks def test_stat(self): @@ -626,13 +614,10 @@ class TestS3FilesStore(unittest.TestCase): ) file_stats = yield store.stat_file("", info=None) - self.assertEqual( - file_stats, - { - "checksum": checksum, - "last_modified": last_modified.timestamp(), - }, - ) + assert file_stats == { + "checksum": checksum, + "last_modified": last_modified.timestamp(), + } stub.assert_no_pending_responses() @@ -655,16 +640,16 @@ class TestGCSFilesStore(unittest.TestCase): expected_policy = {"role": "READER", "entity": "allAuthenticatedUsers"} yield store.persist_file(path, buf, info=None, meta=meta, headers=None) s = yield store.stat_file(path, info=None) - self.assertIn("last_modified", s) - self.assertIn("checksum", s) - self.assertEqual(s["checksum"], "cdcda85605e46d0af6110752770dce3c") + assert "last_modified" in s + assert "checksum" in s + assert s["checksum"] == "cdcda85605e46d0af6110752770dce3c" u = urlparse(uri) content, acl, blob = get_gcs_content_and_delete(u.hostname, u.path[1:] + path) - self.assertEqual(content, data) - self.assertEqual(blob.metadata, {"foo": "bar"}) - self.assertEqual(blob.cache_control, GCSFilesStore.CACHE_CONTROL) - self.assertEqual(blob.content_type, "application/octet-stream") - self.assertIn(expected_policy, acl) + assert content == data + assert blob.metadata == {"foo": "bar"} + assert blob.cache_control == GCSFilesStore.CACHE_CONTROL + assert blob.content_type == "application/octet-stream" + assert expected_policy in acl @defer.inlineCallbacks def test_blob_path_consistency(self): @@ -702,12 +687,12 @@ class TestFTPFileStore(unittest.TestCase): with MockFTPServer() as ftp_server: store = FTPFilesStore(ftp_server.url("/")) empty_dict = yield store.stat_file(path, info=None) - self.assertEqual(empty_dict, {}) + assert empty_dict == {} yield store.persist_file(path, buf, info=None, meta=meta, headers=None) stat = yield store.stat_file(path, info=None) - self.assertIn("last_modified", stat) - self.assertIn("checksum", stat) - self.assertEqual(stat["checksum"], "d113d66b2ec7258724a268bd88eef6b6") + assert "last_modified" in stat + assert "checksum" in stat + assert stat["checksum"] == "d113d66b2ec7258724a268bd88eef6b6" path = f"{store.basedir}/{path}" content = get_ftp_content_and_delete( path, @@ -717,7 +702,7 @@ class TestFTPFileStore(unittest.TestCase): store.password, store.USE_ACTIVE_MODE, ) - self.assertEqual(data, content) + assert data == content class ItemWithFiles(Item): @@ -739,12 +724,12 @@ def _prepare_request_object(item_url, flags=None): # this is separate from the one in test_pipeline_media.py to specifically test FilesPipeline subclasses -class BuildFromCrawlerTestCase(unittest.TestCase): - def setUp(self): +class TestBuildFromCrawler: + def setup_method(self): self.tempdir = mkdtemp() self.crawler = get_crawler(None, {"FILES_STORE": self.tempdir}) - def tearDown(self): + def teardown_method(self): rmtree(self.tempdir) def test_simple(self): @@ -755,7 +740,7 @@ class BuildFromCrawlerTestCase(unittest.TestCase): pipe = Pipeline.from_crawler(self.crawler) assert pipe.crawler == self.crawler assert pipe._fingerprinter - self.assertEqual(len(w), 0) + assert len(w) == 0 assert pipe.store def test_has_old_init(self): @@ -768,7 +753,7 @@ class BuildFromCrawlerTestCase(unittest.TestCase): pipe = Pipeline.from_crawler(self.crawler) assert pipe.crawler == self.crawler assert pipe._fingerprinter - self.assertEqual(len(w), 2) + assert len(w) == 2 assert pipe._init_called def test_has_from_settings(self): @@ -785,7 +770,7 @@ class BuildFromCrawlerTestCase(unittest.TestCase): pipe = Pipeline.from_crawler(self.crawler) assert pipe.crawler == self.crawler assert pipe._fingerprinter - self.assertEqual(len(w), 3) + assert len(w) == 3 assert pipe.store assert pipe._from_settings_called @@ -805,6 +790,6 @@ class BuildFromCrawlerTestCase(unittest.TestCase): pipe = Pipeline.from_crawler(self.crawler) assert pipe.crawler == self.crawler assert pipe._fingerprinter - self.assertEqual(len(w), 0) + assert len(w) == 0 assert pipe.store assert pipe._from_crawler_called diff --git a/tests/test_pipeline_images.py b/tests/test_pipeline_images.py index 1d89e44ce..fef6bbbe9 100644 --- a/tests/test_pipeline_images.py +++ b/tests/test_pipeline_images.py @@ -9,109 +9,106 @@ from tempfile import mkdtemp import attr import pytest from itemadapter import ItemAdapter -from twisted.trial import unittest from scrapy.http import Request, Response from scrapy.item import Field, Item from scrapy.pipelines.images import ImageException, ImagesPipeline from scrapy.utils.test import get_crawler -skip_pillow: str | None try: from PIL import Image except ImportError: - skip_pillow = "Missing Python Imaging Library, install https://pypi.org/pypi/Pillow" + pytest.skip( + "Missing Python Imaging Library, install https://pypi.org/pypi/Pillow", + allow_module_level=True, + ) else: encoders = {"jpeg_encoder", "jpeg_decoder"} if not encoders.issubset(set(Image.core.__dict__)): # type: ignore[attr-defined] - skip_pillow = "Missing JPEG encoders" - else: - skip_pillow = None + pytest.skip("Missing JPEG encoders", allow_module_level=True) -class ImagesPipelineTestCase(unittest.TestCase): - skip = skip_pillow - - def setUp(self): +class TestImagesPipeline: + def setup_method(self): self.tempdir = mkdtemp() crawler = get_crawler() self.pipeline = ImagesPipeline(self.tempdir, crawler=crawler) - def tearDown(self): + def teardown_method(self): rmtree(self.tempdir) def test_file_path(self): file_path = self.pipeline.file_path - self.assertEqual( - file_path(Request("https://dev.mydeco.com/mydeco.gif")), - "full/3fd165099d8e71b8a48b2683946e64dbfad8b52d.jpg", + assert ( + file_path(Request("https://dev.mydeco.com/mydeco.gif")) + == "full/3fd165099d8e71b8a48b2683946e64dbfad8b52d.jpg" ) - self.assertEqual( + assert ( file_path( Request( "http://www.maddiebrown.co.uk///catalogue-items//image_54642_12175_95307.jpg" ) - ), - "full/0ffcd85d563bca45e2f90becd0ca737bc58a00b2.jpg", + ) + == "full/0ffcd85d563bca45e2f90becd0ca737bc58a00b2.jpg" ) - self.assertEqual( + assert ( file_path( Request("https://dev.mydeco.com/two/dirs/with%20spaces%2Bsigns.gif") - ), - "full/b250e3a74fff2e4703e310048a5b13eba79379d2.jpg", + ) + == "full/b250e3a74fff2e4703e310048a5b13eba79379d2.jpg" ) - self.assertEqual( + assert ( file_path( Request( "http://www.dfsonline.co.uk/get_prod_image.php?img=status_0907_mdm.jpg" ) - ), - "full/4507be485f38b0da8a0be9eb2e1dfab8a19223f2.jpg", + ) + == "full/4507be485f38b0da8a0be9eb2e1dfab8a19223f2.jpg" ) - self.assertEqual( - file_path(Request("http://www.dorma.co.uk/images/product_details/2532/")), - "full/97ee6f8a46cbbb418ea91502fd24176865cf39b2.jpg", + assert ( + file_path(Request("http://www.dorma.co.uk/images/product_details/2532/")) + == "full/97ee6f8a46cbbb418ea91502fd24176865cf39b2.jpg" ) - self.assertEqual( - file_path(Request("http://www.dorma.co.uk/images/product_details/2532")), - "full/244e0dd7d96a3b7b01f54eded250c9e272577aa1.jpg", + assert ( + file_path(Request("http://www.dorma.co.uk/images/product_details/2532")) + == "full/244e0dd7d96a3b7b01f54eded250c9e272577aa1.jpg" ) - self.assertEqual( + assert ( file_path( Request("http://www.dorma.co.uk/images/product_details/2532"), response=Response("http://www.dorma.co.uk/images/product_details/2532"), info=object(), - ), - "full/244e0dd7d96a3b7b01f54eded250c9e272577aa1.jpg", + ) + == "full/244e0dd7d96a3b7b01f54eded250c9e272577aa1.jpg" ) def test_thumbnail_name(self): thumb_path = self.pipeline.thumb_path name = "50" - self.assertEqual( - thumb_path(Request("file:///tmp/foo.jpg"), name), - "thumbs/50/38a86208c36e59d4404db9e37ce04be863ef0335.jpg", + assert ( + thumb_path(Request("file:///tmp/foo.jpg"), name) + == "thumbs/50/38a86208c36e59d4404db9e37ce04be863ef0335.jpg" ) - self.assertEqual( - thumb_path(Request("file://foo.png"), name), - "thumbs/50/e55b765eba0ec7348e50a1df496040449071b96a.jpg", + assert ( + thumb_path(Request("file://foo.png"), name) + == "thumbs/50/e55b765eba0ec7348e50a1df496040449071b96a.jpg" ) - self.assertEqual( - thumb_path(Request("file:///tmp/foo"), name), - "thumbs/50/0329ad83ebb8e93ea7c7906d46e9ed55f7349a50.jpg", + assert ( + thumb_path(Request("file:///tmp/foo"), name) + == "thumbs/50/0329ad83ebb8e93ea7c7906d46e9ed55f7349a50.jpg" ) - self.assertEqual( - thumb_path(Request("file:///tmp/some.name/foo"), name), - "thumbs/50/850233df65a5b83361798f532f1fc549cd13cbe9.jpg", + assert ( + thumb_path(Request("file:///tmp/some.name/foo"), name) + == "thumbs/50/850233df65a5b83361798f532f1fc549cd13cbe9.jpg" ) - self.assertEqual( + assert ( thumb_path( Request("file:///tmp/some.name/foo"), name, response=Response("file:///tmp/some.name/foo"), info=object(), - ), - "thumbs/50/850233df65a5b83361798f532f1fc549cd13cbe9.jpg", + ) + == "thumbs/50/850233df65a5b83361798f532f1fc549cd13cbe9.jpg" ) def test_thumbnail_name_from_item(self): @@ -130,8 +127,8 @@ class ImagesPipelineTestCase(unittest.TestCase): ).thumb_path item = {"path": "path-to-store-file"} request = Request("http://example.com") - self.assertEqual( - thumb_path(request, "small", item=item), "thumb/small/path-to-store-file" + assert ( + thumb_path(request, "small", item=item) == "thumb/small/path-to-store-file" ) def test_get_images_exception(self): @@ -169,16 +166,13 @@ class ImagesPipelineTestCase(unittest.TestCase): ) path, new_im, new_buf = next(get_images_gen) - self.assertEqual(path, "full/3fd165099d8e71b8a48b2683946e64dbfad8b52d.jpg") - self.assertEqual(orig_im, new_im) - self.assertEqual(buf.getvalue(), new_buf.getvalue()) + assert path == "full/3fd165099d8e71b8a48b2683946e64dbfad8b52d.jpg" + assert orig_im == new_im + assert buf.getvalue() == new_buf.getvalue() thumb_path, thumb_img, thumb_buf = next(get_images_gen) - self.assertEqual( - thumb_path, "thumbs/small/3fd165099d8e71b8a48b2683946e64dbfad8b52d.jpg" - ) - self.assertEqual(thumb_img, thumb_img) - self.assertEqual(orig_thumb_buf.getvalue(), thumb_buf.getvalue()) + assert thumb_path == "thumbs/small/3fd165099d8e71b8a48b2683946e64dbfad8b52d.jpg" + assert orig_thumb_buf.getvalue() == thumb_buf.getvalue() def test_convert_image(self): SIZE = (100, 100) @@ -186,37 +180,35 @@ class ImagesPipelineTestCase(unittest.TestCase): COLOUR = (0, 127, 255) im, buf = _create_image("JPEG", "RGB", SIZE, COLOUR) converted, converted_buf = self.pipeline.convert_image(im, response_body=buf) - self.assertEqual(converted.mode, "RGB") - self.assertEqual(converted.getcolors(), [(10000, COLOUR)]) + assert converted.mode == "RGB" + assert converted.getcolors() == [(10000, COLOUR)] # check that we don't convert JPEGs again - self.assertEqual(converted_buf, buf) + assert converted_buf == buf # check that thumbnail keep image ratio thumbnail, _ = self.pipeline.convert_image( converted, size=(10, 25), response_body=converted_buf ) - self.assertEqual(thumbnail.mode, "RGB") - self.assertEqual(thumbnail.size, (10, 10)) + assert thumbnail.mode == "RGB" + assert thumbnail.size == (10, 10) # transparency case: RGBA and PNG COLOUR = (0, 127, 255, 50) im, buf = _create_image("PNG", "RGBA", SIZE, COLOUR) converted, _ = self.pipeline.convert_image(im, response_body=buf) - self.assertEqual(converted.mode, "RGB") - self.assertEqual(converted.getcolors(), [(10000, (205, 230, 255))]) + assert converted.mode == "RGB" + assert converted.getcolors() == [(10000, (205, 230, 255))] # transparency case with palette: P and PNG COLOUR = (0, 127, 255, 50) im, buf = _create_image("PNG", "RGBA", SIZE, COLOUR) im = im.convert("P") converted, _ = self.pipeline.convert_image(im, response_body=buf) - self.assertEqual(converted.mode, "RGB") - self.assertEqual(converted.getcolors(), [(10000, (205, 230, 255))]) + assert converted.mode == "RGB" + assert converted.getcolors() == [(10000, (205, 230, 255))] class ImagesPipelineTestCaseFieldsMixin: - skip = skip_pillow - def test_item_fields_default(self): url = "http://www.example.com/images/1.jpg" item = self.item_class(name="item1", image_urls=[url]) @@ -224,12 +216,12 @@ class ImagesPipelineTestCaseFieldsMixin: get_crawler(None, {"IMAGES_STORE": "s3://example/images/"}) ) requests = list(pipeline.get_media_requests(item, None)) - self.assertEqual(requests[0].url, url) + assert requests[0].url == url results = [(True, {"url": url})] item = pipeline.item_completed(results, item, None) images = ItemAdapter(item).get("images") - self.assertEqual(images, [results[0][1]]) - self.assertIsInstance(item, self.item_class) + assert images == [results[0][1]] + assert isinstance(item, self.item_class) def test_item_fields_override_settings(self): url = "http://www.example.com/images/1.jpg" @@ -245,17 +237,15 @@ class ImagesPipelineTestCaseFieldsMixin: ) ) requests = list(pipeline.get_media_requests(item, None)) - self.assertEqual(requests[0].url, url) + assert requests[0].url == url results = [(True, {"url": url})] item = pipeline.item_completed(results, item, None) custom_images = ItemAdapter(item).get("custom_images") - self.assertEqual(custom_images, [results[0][1]]) - self.assertIsInstance(item, self.item_class) + assert custom_images == [results[0][1]] + assert isinstance(item, self.item_class) -class ImagesPipelineTestCaseFieldsDict( - ImagesPipelineTestCaseFieldsMixin, unittest.TestCase -): +class TestImagesPipelineFieldsDict(ImagesPipelineTestCaseFieldsMixin): item_class = dict @@ -269,9 +259,7 @@ class ImagesPipelineTestItem(Item): custom_images = Field() -class ImagesPipelineTestCaseFieldsItem( - ImagesPipelineTestCaseFieldsMixin, unittest.TestCase -): +class TestImagesPipelineFieldsItem(ImagesPipelineTestCaseFieldsMixin): item_class = ImagesPipelineTestItem @@ -286,9 +274,7 @@ class ImagesPipelineTestDataClass: custom_images: list = dataclasses.field(default_factory=list) -class ImagesPipelineTestCaseFieldsDataClass( - ImagesPipelineTestCaseFieldsMixin, unittest.TestCase -): +class TestImagesPipelineFieldsDataClass(ImagesPipelineTestCaseFieldsMixin): item_class = ImagesPipelineTestDataClass @@ -303,15 +289,11 @@ class ImagesPipelineTestAttrsItem: custom_images: list[dict[str, str]] = attr.ib(default=list) -class ImagesPipelineTestCaseFieldsAttrsItem( - ImagesPipelineTestCaseFieldsMixin, unittest.TestCase -): +class TestImagesPipelineFieldsAttrsItem(ImagesPipelineTestCaseFieldsMixin): item_class = ImagesPipelineTestAttrsItem -class ImagesPipelineTestCaseCustomSettings(unittest.TestCase): - skip = skip_pillow - +class TestImagesPipelineCustomSettings: img_cls_attribute_names = [ # Pipeline attribute names with corresponding setting names. ("EXPIRES", "IMAGES_EXPIRES"), @@ -332,10 +314,10 @@ class ImagesPipelineTestCaseCustomSettings(unittest.TestCase): "IMAGES_RESULT_FIELD": "images", } - def setUp(self): + def setup_method(self): self.tempdir = mkdtemp() - def tearDown(self): + def teardown_method(self): rmtree(self.tempdir) def _generate_fake_settings(self, prefix=None): @@ -397,11 +379,11 @@ class ImagesPipelineTestCaseCustomSettings(unittest.TestCase): for pipe_attr, settings_attr in self.img_cls_attribute_names: expected_default_value = self.default_pipeline_settings.get(pipe_attr) custom_value = custom_settings.get(settings_attr) - self.assertNotEqual(expected_default_value, custom_value) - self.assertEqual( - getattr(default_sts_pipe, pipe_attr.lower()), expected_default_value + assert expected_default_value != custom_value + assert ( + getattr(default_sts_pipe, pipe_attr.lower()) == expected_default_value ) - self.assertEqual(getattr(user_sts_pipe, pipe_attr.lower()), custom_value) + assert getattr(user_sts_pipe, pipe_attr.lower()) == custom_value def test_subclass_attrs_preserved_default_settings(self): """ @@ -415,8 +397,8 @@ class ImagesPipelineTestCaseCustomSettings(unittest.TestCase): for pipe_attr, settings_attr in self.img_cls_attribute_names: # Instance attribute (lowercase) must be equal to class attribute (uppercase). attr_value = getattr(pipeline, pipe_attr.lower()) - self.assertNotEqual(attr_value, self.default_pipeline_settings[pipe_attr]) - self.assertEqual(attr_value, getattr(pipeline, pipe_attr)) + assert attr_value != self.default_pipeline_settings[pipe_attr] + assert attr_value == getattr(pipeline, pipe_attr) def test_subclass_attrs_preserved_custom_settings(self): """ @@ -430,9 +412,9 @@ class ImagesPipelineTestCaseCustomSettings(unittest.TestCase): # Instance attribute (lowercase) must be equal to # value defined in settings. value = getattr(pipeline, pipe_attr.lower()) - self.assertNotEqual(value, self.default_pipeline_settings[pipe_attr]) + assert value != self.default_pipeline_settings[pipe_attr] setings_value = settings.get(settings_attr) - self.assertEqual(value, setings_value) + assert value == setings_value def test_no_custom_settings_for_subclasses(self): """ @@ -449,7 +431,7 @@ class ImagesPipelineTestCaseCustomSettings(unittest.TestCase): for pipe_attr, settings_attr in self.img_cls_attribute_names: # Values from settings for custom pipeline should be set on pipeline instance. custom_value = self.default_pipeline_settings.get(pipe_attr.upper()) - self.assertEqual(getattr(user_pipeline, pipe_attr.lower()), custom_value) + assert getattr(user_pipeline, pipe_attr.lower()) == custom_value def test_custom_settings_for_subclasses(self): """ @@ -468,8 +450,8 @@ class ImagesPipelineTestCaseCustomSettings(unittest.TestCase): for pipe_attr, settings_attr in self.img_cls_attribute_names: # Values from settings for custom pipeline should be set on pipeline instance. custom_value = settings.get(prefix + "_" + settings_attr) - self.assertNotEqual(custom_value, self.default_pipeline_settings[pipe_attr]) - self.assertEqual(getattr(user_pipeline, pipe_attr.lower()), custom_value) + assert custom_value != self.default_pipeline_settings[pipe_attr] + assert getattr(user_pipeline, pipe_attr.lower()) == custom_value def test_custom_settings_and_class_attrs_for_subclasses(self): """ @@ -482,8 +464,8 @@ class ImagesPipelineTestCaseCustomSettings(unittest.TestCase): user_pipeline = pipeline_cls.from_crawler(get_crawler(None, settings)) for pipe_attr, settings_attr in self.img_cls_attribute_names: custom_value = settings.get(prefix + "_" + settings_attr) - self.assertNotEqual(custom_value, self.default_pipeline_settings[pipe_attr]) - self.assertEqual(getattr(user_pipeline, pipe_attr.lower()), custom_value) + assert custom_value != self.default_pipeline_settings[pipe_attr] + assert getattr(user_pipeline, pipe_attr.lower()) == custom_value def test_cls_attrs_with_DEFAULT_prefix(self): class UserDefinedImagePipeline(ImagesPipeline): @@ -493,13 +475,13 @@ class ImagesPipelineTestCaseCustomSettings(unittest.TestCase): pipeline = UserDefinedImagePipeline.from_crawler( get_crawler(None, {"IMAGES_STORE": self.tempdir}) ) - self.assertEqual( - pipeline.images_result_field, - UserDefinedImagePipeline.DEFAULT_IMAGES_RESULT_FIELD, + assert ( + pipeline.images_result_field + == UserDefinedImagePipeline.DEFAULT_IMAGES_RESULT_FIELD ) - self.assertEqual( - pipeline.images_urls_field, - UserDefinedImagePipeline.DEFAULT_IMAGES_URLS_FIELD, + assert ( + pipeline.images_urls_field + == UserDefinedImagePipeline.DEFAULT_IMAGES_URLS_FIELD ) def test_user_defined_subclass_default_key_names(self): @@ -516,7 +498,7 @@ class ImagesPipelineTestCaseCustomSettings(unittest.TestCase): for pipe_attr, settings_attr in self.img_cls_attribute_names: expected_value = settings.get(settings_attr) - self.assertEqual(getattr(pipeline_cls, pipe_attr.lower()), expected_value) + assert getattr(pipeline_cls, pipe_attr.lower()) == expected_value def _create_image(format, *a, **kw): diff --git a/tests/test_pipeline_media.py b/tests/test_pipeline_media.py index c6fdd3767..d915fc2a3 100644 --- a/tests/test_pipeline_media.py +++ b/tests/test_pipeline_media.py @@ -2,6 +2,7 @@ from __future__ import annotations import warnings +import pytest from testfixtures import LogCapture from twisted.internet import reactor from twisted.internet.defer import Deferred, inlineCallbacks @@ -18,15 +19,6 @@ from scrapy.utils.log import failure_to_exc_info from scrapy.utils.signal import disconnect_all from scrapy.utils.test import get_crawler -try: - from PIL import Image # noqa: F401 -except ImportError: - skip_pillow: str | None = ( - "Missing Python Imaging Library, install https://pypi.org/pypi/Pillow" - ) -else: - skip_pillow = None - def _mocked_download_func(request, info): assert request.callback is NO_CALLBACK @@ -51,7 +43,7 @@ class UserDefinedPipeline(MediaPipeline): return "" -class BaseMediaPipelineTestCase(unittest.TestCase): +class TestBaseMediaPipeline(unittest.TestCase): pipeline_class = UserDefinedPipeline settings = None @@ -123,9 +115,9 @@ class BaseMediaPipelineTestCase(unittest.TestCase): failure = Failure(file_exc) # The Failure should encapsulate a FileException ... - self.assertEqual(failure.value, file_exc) + assert failure.value == file_exc # ... and it should have the StopIteration exception set as its context - self.assertEqual(failure.value.__context__, def_gen_return_exc) + assert failure.value.__context__ == def_gen_return_exc # Let's calculate the request fingerprint and fake some runtime data... fp = self.fingerprint(request) @@ -136,12 +128,12 @@ class BaseMediaPipelineTestCase(unittest.TestCase): # When calling the method that caches the Request's result ... self.pipe._cache_result_and_execute_waiters(failure, fp, info) # ... it should store the Twisted Failure ... - self.assertEqual(info.downloaded[fp], failure) + assert info.downloaded[fp] == failure # ... encapsulating the original FileException ... - self.assertEqual(info.downloaded[fp].value, file_exc) + assert info.downloaded[fp].value == file_exc # ... but it should not store the StopIteration exception on its context context = getattr(info.downloaded[fp].value, "__context__", None) - self.assertIsNone(context) + assert context is None def test_default_item_completed(self): item = {"name": "name"} @@ -158,7 +150,7 @@ class BaseMediaPipelineTestCase(unittest.TestCase): assert len(log.records) == 1 record = log.records[0] assert record.levelname == "ERROR" - self.assertTupleEqual(record.exc_info, failure_to_exc_info(fail)) + assert record.exc_info == failure_to_exc_info(fail) # disable failure logging and check again self.pipe.LOG_FAILED_RESULTS = False @@ -208,7 +200,7 @@ class MockedMediaPipeline(UserDefinedPipeline): return item -class MediaPipelineTestCase(BaseMediaPipelineTestCase): +class TestMediaPipeline(TestBaseMediaPipeline): pipeline_class = MockedMediaPipeline def _errback(self, result): @@ -225,16 +217,13 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase): ) item = {"requests": req} new_item = yield self.pipe.process_item(item, self.spider) - self.assertEqual(new_item["results"], [(True, {})]) - self.assertEqual( - self.pipe._mockcalled, - [ - "get_media_requests", - "media_to_download", - "media_downloaded", - "item_completed", - ], - ) + assert new_item["results"] == [(True, {})] + assert self.pipe._mockcalled == [ + "get_media_requests", + "media_to_download", + "media_downloaded", + "item_completed", + ] @inlineCallbacks def test_result_failure(self): @@ -247,17 +236,14 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase): ) item = {"requests": req} new_item = yield self.pipe.process_item(item, self.spider) - self.assertEqual(new_item["results"], [(False, fail)]) - self.assertEqual( - self.pipe._mockcalled, - [ - "get_media_requests", - "media_to_download", - "media_failed", - "request_errback", - "item_completed", - ], - ) + assert new_item["results"] == [(False, fail)] + assert self.pipe._mockcalled == [ + "get_media_requests", + "media_to_download", + "media_failed", + "request_errback", + "item_completed", + ] @inlineCallbacks def test_mix_of_success_and_failure(self): @@ -268,18 +254,18 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase): req2 = Request("http://url2", meta={"response": fail}) item = {"requests": [req1, req2]} new_item = yield self.pipe.process_item(item, self.spider) - self.assertEqual(new_item["results"], [(True, {}), (False, fail)]) + assert new_item["results"] == [(True, {}), (False, fail)] m = self.pipe._mockcalled # only once - self.assertEqual(m[0], "get_media_requests") # first hook called - self.assertEqual(m.count("get_media_requests"), 1) - self.assertEqual(m.count("item_completed"), 1) - self.assertEqual(m[-1], "item_completed") # last hook called + assert m[0] == "get_media_requests" # first hook called + assert m.count("get_media_requests") == 1 + assert m.count("item_completed") == 1 + assert m[-1] == "item_completed" # last hook called # twice, one per request - self.assertEqual(m.count("media_to_download"), 2) + assert m.count("media_to_download") == 2 # one to handle success and other for failure - self.assertEqual(m.count("media_downloaded"), 1) - self.assertEqual(m.count("media_failed"), 1) + assert m.count("media_downloaded") == 1 + assert m.count("media_failed") == 1 @inlineCallbacks def test_get_media_requests(self): @@ -288,7 +274,7 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase): item = {"requests": req} # pass a single item new_item = yield self.pipe.process_item(item, self.spider) assert new_item is item - self.assertIn(self.fingerprint(req), self.info.downloaded) + assert self.fingerprint(req) in self.info.downloaded # returns iterable of Requests req1 = Request("http://url1") @@ -305,8 +291,8 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase): req1 = Request("http://url1", meta={"response": rsp1}) item = {"requests": req1} new_item = yield self.pipe.process_item(item, self.spider) - self.assertTrue(new_item is item) - self.assertEqual(new_item["results"], [(True, {})]) + assert new_item is item + assert new_item["results"] == [(True, {})] # rsp2 is ignored, rsp1 must be in results because request fingerprints are the same req2 = Request( @@ -314,9 +300,9 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase): ) item = {"requests": req2} new_item = yield self.pipe.process_item(item, self.spider) - self.assertTrue(new_item is item) - self.assertEqual(self.fingerprint(req1), self.fingerprint(req2)) - self.assertEqual(new_item["results"], [(True, {})]) + assert new_item is item + assert self.fingerprint(req1) == self.fingerprint(req2) + assert new_item["results"] == [(True, {})] @inlineCallbacks def test_results_are_cached_for_requests_of_single_item(self): @@ -327,17 +313,17 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase): ) item = {"requests": [req1, req2]} new_item = yield self.pipe.process_item(item, self.spider) - self.assertTrue(new_item is item) - self.assertEqual(new_item["results"], [(True, {}), (True, {})]) + assert new_item is item + assert new_item["results"] == [(True, {}), (True, {})] @inlineCallbacks def test_wait_if_request_is_downloading(self): def _check_downloading(response): fp = self.fingerprint(req1) - self.assertTrue(fp in self.info.downloading) - self.assertTrue(fp in self.info.waiting) - self.assertTrue(fp not in self.info.downloaded) - self.assertEqual(len(self.info.waiting[fp]), 2) + assert fp in self.info.downloading + assert fp in self.info.waiting + assert fp not in self.info.downloaded + assert len(self.info.waiting[fp]) == 2 return response rsp1 = Response("http://url") @@ -348,39 +334,40 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase): return dfd def rsp2_func(): - self.fail("it must cache rsp1 result and must not try to redownload") + pytest.fail("it must cache rsp1 result and must not try to redownload") req1 = Request("http://url", meta={"response": rsp1_func}) req2 = Request(req1.url, meta={"response": rsp2_func}) item = {"requests": [req1, req2]} new_item = yield self.pipe.process_item(item, self.spider) - self.assertEqual(new_item["results"], [(True, {}), (True, {})]) + assert new_item["results"] == [(True, {}), (True, {})] @inlineCallbacks def test_use_media_to_download_result(self): req = Request("http://url", meta={"result": "ITSME", "response": self.fail}) item = {"requests": req} new_item = yield self.pipe.process_item(item, self.spider) - self.assertEqual(new_item["results"], [(True, "ITSME")]) - self.assertEqual( - self.pipe._mockcalled, - ["get_media_requests", "media_to_download", "item_completed"], - ) + assert new_item["results"] == [(True, "ITSME")] + assert self.pipe._mockcalled == [ + "get_media_requests", + "media_to_download", + "item_completed", + ] def test_key_for_pipe(self): - self.assertEqual( - self.pipe._key_for_pipe("IMAGES", base_class_name="MediaPipeline"), - "MOCKEDMEDIAPIPELINE_IMAGES", + assert ( + self.pipe._key_for_pipe("IMAGES", base_class_name="MediaPipeline") + == "MOCKEDMEDIAPIPELINE_IMAGES" ) -class MediaPipelineAllowRedirectSettingsTestCase(unittest.TestCase): +class TestMediaPipelineAllowRedirectSettings: def _assert_request_no3xx(self, pipeline_class, settings): pipe = pipeline_class(crawler=get_crawler(None, settings)) request = Request("http://url") pipe._modify_media_request(request) - self.assertIn("handle_httpstatus_list", request.meta) + assert "handle_httpstatus_list" in request.meta for status, check in [ (200, True), # These are the status codes we want @@ -396,9 +383,9 @@ class MediaPipelineAllowRedirectSettingsTestCase(unittest.TestCase): (500, True), ]: if check: - self.assertIn(status, request.meta["handle_httpstatus_list"]) + assert status in request.meta["handle_httpstatus_list"] else: - self.assertNotIn(status, request.meta["handle_httpstatus_list"]) + assert status not in request.meta["handle_httpstatus_list"] def test_subclass_standard_setting(self): self._assert_request_no3xx(UserDefinedPipeline, {"MEDIA_ALLOW_REDIRECTS": True}) @@ -409,8 +396,8 @@ class MediaPipelineAllowRedirectSettingsTestCase(unittest.TestCase): ) -class BuildFromCrawlerTestCase(unittest.TestCase): - def setUp(self): +class TestBuildFromCrawler: + def setup_method(self): self.crawler = get_crawler(None, {"FILES_STORE": "/foo"}) def test_simple(self): @@ -421,7 +408,7 @@ class BuildFromCrawlerTestCase(unittest.TestCase): pipe = Pipeline.from_crawler(self.crawler) assert pipe.crawler == self.crawler assert pipe._fingerprinter - self.assertEqual(len(w), 0) + assert len(w) == 0 def test_has_old_init(self): class Pipeline(UserDefinedPipeline): @@ -433,7 +420,7 @@ class BuildFromCrawlerTestCase(unittest.TestCase): pipe = Pipeline.from_crawler(self.crawler) assert pipe.crawler == self.crawler assert pipe._fingerprinter - self.assertEqual(len(w), 2) + assert len(w) == 2 assert pipe._init_called def test_has_from_settings(self): @@ -450,7 +437,7 @@ class BuildFromCrawlerTestCase(unittest.TestCase): pipe = Pipeline.from_crawler(self.crawler) assert pipe.crawler == self.crawler assert pipe._fingerprinter - self.assertEqual(len(w), 2) + assert len(w) == 2 assert pipe._from_settings_called def test_has_from_settings_and_from_crawler(self): @@ -474,7 +461,7 @@ class BuildFromCrawlerTestCase(unittest.TestCase): pipe = Pipeline.from_crawler(self.crawler) assert pipe.crawler == self.crawler assert pipe._fingerprinter - self.assertEqual(len(w), 2) + assert len(w) == 2 assert pipe._from_settings_called assert pipe._from_crawler_called @@ -497,7 +484,7 @@ class BuildFromCrawlerTestCase(unittest.TestCase): pipe = Pipeline.from_crawler(self.crawler) assert pipe.crawler == self.crawler assert pipe._fingerprinter - self.assertEqual(len(w), 2) + assert len(w) == 2 assert pipe._from_settings_called assert pipe._init_called @@ -521,7 +508,7 @@ class BuildFromCrawlerTestCase(unittest.TestCase): pipe = Pipeline.from_crawler(self.crawler) assert pipe.crawler == self.crawler assert pipe._fingerprinter - self.assertEqual(len(w), 0) + assert len(w) == 0 assert pipe._from_crawler_called assert pipe._init_called @@ -542,5 +529,5 @@ class BuildFromCrawlerTestCase(unittest.TestCase): # this and the next assert will fail as MediaPipeline.from_crawler() wasn't called assert pipe.crawler == self.crawler assert pipe._fingerprinter - self.assertEqual(len(w), 0) + assert len(w) == 0 assert pipe._from_crawler_called From 380c2279b92f1aa7386e79fc43109499a057e8cf Mon Sep 17 00:00:00 2001 From: Andrey Rakhmatullin <wrar@wrar.name> Date: Sun, 9 Mar 2025 23:23:51 +0400 Subject: [PATCH 3/5] Converting tests to plain asserts, part 7. (#6710) --- tests/test_downloader_handlers.py | 243 ++++++++-------- tests/test_downloader_handlers_http2.py | 41 ++- tests/test_exporters.py | 168 +++++------ tests/test_feedexport.py | 367 +++++++++++------------- tests/test_http2_client_protocol.py | 123 ++++---- tests/test_http_cookies.py | 52 ++-- tests/test_http_headers.py | 87 +++--- 7 files changed, 512 insertions(+), 569 deletions(-) diff --git a/tests/test_downloader_handlers.py b/tests/test_downloader_handlers.py index 323a51002..19bd02498 100644 --- a/tests/test_downloader_handlers.py +++ b/tests/test_downloader_handlers.py @@ -69,46 +69,46 @@ class OffDH: return cls(crawler) -class LoadTestCase(unittest.TestCase): +class TestLoad: def test_enabled_handler(self): handlers = {"scheme": DummyDH} crawler = get_crawler(settings_dict={"DOWNLOAD_HANDLERS": handlers}) dh = DownloadHandlers(crawler) - self.assertIn("scheme", dh._schemes) - self.assertIn("scheme", dh._handlers) - self.assertNotIn("scheme", dh._notconfigured) + assert "scheme" in dh._schemes + assert "scheme" in dh._handlers + assert "scheme" not in dh._notconfigured def test_not_configured_handler(self): handlers = {"scheme": OffDH} crawler = get_crawler(settings_dict={"DOWNLOAD_HANDLERS": handlers}) dh = DownloadHandlers(crawler) - self.assertIn("scheme", dh._schemes) - self.assertNotIn("scheme", dh._handlers) - self.assertIn("scheme", dh._notconfigured) + assert "scheme" in dh._schemes + assert "scheme" not in dh._handlers + assert "scheme" in dh._notconfigured def test_disabled_handler(self): handlers = {"scheme": None} crawler = get_crawler(settings_dict={"DOWNLOAD_HANDLERS": handlers}) dh = DownloadHandlers(crawler) - self.assertNotIn("scheme", dh._schemes) + assert "scheme" not in dh._schemes for scheme in handlers: # force load handlers dh._get_handler(scheme) - self.assertNotIn("scheme", dh._handlers) - self.assertIn("scheme", dh._notconfigured) + assert "scheme" not in dh._handlers + assert "scheme" in dh._notconfigured def test_lazy_handlers(self): handlers = {"scheme": DummyLazyDH} crawler = get_crawler(settings_dict={"DOWNLOAD_HANDLERS": handlers}) dh = DownloadHandlers(crawler) - self.assertIn("scheme", dh._schemes) - self.assertNotIn("scheme", dh._handlers) + assert "scheme" in dh._schemes + assert "scheme" not in dh._handlers for scheme in handlers: # force load lazy handler dh._get_handler(scheme) - self.assertIn("scheme", dh._handlers) - self.assertNotIn("scheme", dh._notconfigured) + assert "scheme" in dh._handlers + assert "scheme" not in dh._notconfigured -class FileTestCase(unittest.TestCase): +class TestFile(unittest.TestCase): def setUp(self): # add a special char to check that they are handled correctly self.fd, self.tmpname = mkstemp(suffix="^") @@ -122,10 +122,10 @@ class FileTestCase(unittest.TestCase): def test_download(self): def _test(response): - self.assertEqual(response.url, request.url) - self.assertEqual(response.status, 200) - self.assertEqual(response.body, b"0123456789") - self.assertEqual(response.protocol, None) + assert response.url == request.url + assert response.status == 200 + assert response.body == b"0123456789" + assert response.protocol is None request = Request(path_to_file_uri(self.tmpname)) assert request.url.upper().endswith("%5E") @@ -217,7 +217,7 @@ class DuplicateHeaderResource(resource.Resource): return b"" -class HttpTestCase(unittest.TestCase, ABC): +class TestHttp(unittest.TestCase, ABC): scheme = "http" # only used for HTTPS tests @@ -336,8 +336,8 @@ class HttpTestCase(unittest.TestCase, ABC): def test_host_header_not_in_request_headers(self): def _test(response): - self.assertEqual(response.body, to_bytes(f"{self.host}:{self.portno}")) - self.assertEqual(request.headers, {}) + assert response.body == to_bytes(f"{self.host}:{self.portno}") + assert not request.headers request = Request(self.getURL("host")) return self.download_request(request, Spider("foo")).addCallback(_test) @@ -346,8 +346,8 @@ class HttpTestCase(unittest.TestCase, ABC): host = self.host + ":" + str(self.portno) def _test(response): - self.assertEqual(response.body, host.encode()) - self.assertEqual(request.headers.get("Host"), host.encode()) + assert response.body == host.encode() + assert request.headers.get("Host") == host.encode() request = Request(self.getURL("host"), headers={"Host": host}) return self.download_request(request, Spider("foo")).addCallback(_test) @@ -365,7 +365,7 @@ class HttpTestCase(unittest.TestCase, ABC): """ def _test(response): - self.assertEqual(response.body, b"0") + assert response.body == b"0" request = Request(self.getURL("contentlength"), method="POST") return self.download_request(request, Spider("foo")).addCallback(_test) @@ -376,8 +376,8 @@ class HttpTestCase(unittest.TestCase, ABC): headers = Headers(json.loads(response.text)["headers"]) contentlengths = headers.getlist("Content-Length") - self.assertEqual(len(contentlengths), 1) - self.assertEqual(contentlengths, [b"0"]) + assert len(contentlengths) == 1 + assert contentlengths == [b"0"] request = Request(self.getURL("echo"), method="POST") return self.download_request(request, Spider("foo")).addCallback(_test) @@ -399,7 +399,7 @@ class HttpTestCase(unittest.TestCase, ABC): def _test_response_class(self, filename, body, response_class): def _test(response): - self.assertEqual(type(response), response_class) + assert type(response) is response_class # pylint: disable=unidiomatic-typecheck request = Request(self.getURL(filename), body=body) return self.download_request(request, Spider("foo")).addCallback(_test) @@ -416,17 +416,14 @@ class HttpTestCase(unittest.TestCase, ABC): def test_get_duplicate_header(self): def _test(response): - self.assertEqual( - response.headers.getlist(b"Set-Cookie"), - [b"a=b", b"c=d"], - ) + assert response.headers.getlist(b"Set-Cookie") == [b"a=b", b"c=d"] request = Request(self.getURL("duplicate-header")) return self.download_request(request, Spider("foo")).addCallback(_test) @pytest.mark.filterwarnings("ignore::scrapy.exceptions.ScrapyDeprecationWarning") -class Http10TestCase(HttpTestCase): +class TestHttp10(TestHttp): """HTTP 1.0 test case""" @property @@ -441,11 +438,11 @@ class Http10TestCase(HttpTestCase): return d -class Https10TestCase(Http10TestCase): +class TestHttps10(TestHttp10): scheme = "https" -class Http11TestCase(HttpTestCase): +class TestHttp11(TestHttp): """HTTP 1.1 test case""" @property @@ -466,7 +463,7 @@ class Http11TestCase(HttpTestCase): body = b"Some plain text\ndata with tabs\t and null bytes\0" def _test_type(response): - self.assertEqual(type(response), TextResponse) + assert type(response) is TextResponse # pylint: disable=unidiomatic-typecheck request = Request(self.getURL("nocontenttype"), body=body) d = self.download_request(request, Spider("foo")) @@ -583,7 +580,7 @@ class Http11TestCase(HttpTestCase): return d -class Https11TestCase(Http11TestCase): +class TestHttps11(TestHttp11): scheme = "https" tls_log_message = ( @@ -611,7 +608,7 @@ class Https11TestCase(Http11TestCase): yield download_handler.close() -class SimpleHttpsTest(unittest.TestCase): +class TestSimpleHttps(unittest.TestCase): """Base class for special cases tested with just one simple request""" keyfile = "keys/localhost.key" @@ -663,7 +660,7 @@ class SimpleHttpsTest(unittest.TestCase): return d -class Https11WrongHostnameTestCase(SimpleHttpsTest): +class TestHttps11WrongHostname(TestSimpleHttps): # above tests use a server certificate for "localhost", # client connection to "localhost" too. # here we test that even if the server certificate is for another domain, @@ -673,7 +670,7 @@ class Https11WrongHostnameTestCase(SimpleHttpsTest): certfile = "keys/example-com.cert.pem" -class Https11InvalidDNSId(SimpleHttpsTest): +class TestHttps11InvalidDNSId(TestSimpleHttps): """Connect to HTTPS hosts with IP while certificate uses domain names IDs.""" def setUp(self): @@ -681,18 +678,18 @@ class Https11InvalidDNSId(SimpleHttpsTest): self.host = "127.0.0.1" -class Https11InvalidDNSPattern(SimpleHttpsTest): +class TestHttps11InvalidDNSPattern(TestSimpleHttps): """Connect to HTTPS hosts where the certificate are issued to an ip instead of a domain.""" keyfile = "keys/localhost.ip.key" certfile = "keys/localhost.ip.crt" -class Https11CustomCiphers(SimpleHttpsTest): +class TestHttps11CustomCiphers(TestSimpleHttps): cipher_string = "CAMELLIA256-SHA" -class Http11MockServerTestCase(unittest.TestCase): +class TestHttp11MockServer(unittest.TestCase): """HTTP 1.1 test case with MockServer""" settings_dict: dict | None = None @@ -719,7 +716,7 @@ class Http11MockServerTestCase(unittest.TestCase): ) ) failure = crawler.spider.meta["failure"] - self.assertIsInstance(failure.value, defer.CancelledError) + assert isinstance(failure.value, defer.CancelledError) @defer.inlineCallbacks def test_download(self): @@ -728,9 +725,9 @@ class Http11MockServerTestCase(unittest.TestCase): seed=Request(url=self.mockserver.url("", is_secure=self.is_secure)) ) failure = crawler.spider.meta.get("failure") - self.assertTrue(failure is None) + assert failure is None reason = crawler.spider.meta["close_reason"] - self.assertTrue(reason, "finished") + assert reason == "finished" class UriResource(resource.Resource): @@ -748,7 +745,7 @@ class UriResource(resource.Resource): return b"" -class HttpProxyTestCase(unittest.TestCase, ABC): +class TestHttpProxy(unittest.TestCase, ABC): expected_http_proxy_request_body = b"http://example.com" @property @@ -777,9 +774,9 @@ class HttpProxyTestCase(unittest.TestCase, ABC): def test_download_with_proxy(self): def _test(response): - self.assertEqual(response.status, 200) - self.assertEqual(response.url, request.url) - self.assertEqual(response.body, self.expected_http_proxy_request_body) + assert response.status == 200 + assert response.url == request.url + assert response.body == self.expected_http_proxy_request_body http_proxy = self.getURL("") request = Request("http://example.com", meta={"proxy": http_proxy}) @@ -787,22 +784,22 @@ class HttpProxyTestCase(unittest.TestCase, ABC): def test_download_without_proxy(self): def _test(response): - self.assertEqual(response.status, 200) - self.assertEqual(response.url, request.url) - self.assertEqual(response.body, b"/path/to/resource") + assert response.status == 200 + assert response.url == request.url + assert response.body == b"/path/to/resource" request = Request(self.getURL("path/to/resource")) return self.download_request(request, Spider("foo")).addCallback(_test) @pytest.mark.filterwarnings("ignore::scrapy.exceptions.ScrapyDeprecationWarning") -class Http10ProxyTestCase(HttpProxyTestCase): +class TestHttp10Proxy(TestHttpProxy): @property def download_handler_cls(self) -> type[DownloadHandlerProtocol]: return HTTP10DownloadHandler -class Http11ProxyTestCase(HttpProxyTestCase): +class TestHttp11Proxy(TestHttpProxy): @property def download_handler_cls(self) -> type[DownloadHandlerProtocol]: return HTTP11DownloadHandler @@ -817,13 +814,13 @@ class Http11ProxyTestCase(HttpProxyTestCase): request = Request(domain, meta={"proxy": http_proxy, "download_timeout": 0.2}) d = self.download_request(request, Spider("foo")) timeout = yield self.assertFailure(d, error.TimeoutError) - self.assertIn(domain, timeout.osError) + assert domain in timeout.osError def test_download_with_proxy_without_http_scheme(self): def _test(response): - self.assertEqual(response.status, 200) - self.assertEqual(response.url, request.url) - self.assertEqual(response.body, self.expected_http_proxy_request_body) + assert response.status == 200 + assert response.url == request.url + assert response.body == self.expected_http_proxy_request_body http_proxy = self.getURL("").replace("http://", "") request = Request("http://example.com", meta={"proxy": http_proxy}) @@ -839,8 +836,8 @@ class HttpDownloadHandlerMock: @pytest.mark.requires_botocore -class S3AnonTestCase(unittest.TestCase): - def setUp(self): +class TestS3Anon: + def setup_method(self): crawler = get_crawler() self.s3reqh = build_from_crawler( S3DownloadHandler, @@ -854,13 +851,13 @@ class S3AnonTestCase(unittest.TestCase): def test_anon_request(self): req = Request("s3://aws-publicdatasets/") httpreq = self.download_request(req, self.spider) - self.assertEqual(hasattr(self.s3reqh, "anon"), True) - self.assertEqual(self.s3reqh.anon, True) - self.assertEqual(httpreq.url, "http://aws-publicdatasets.s3.amazonaws.com/") + assert hasattr(self.s3reqh, "anon") + assert self.s3reqh.anon + assert httpreq.url == "http://aws-publicdatasets.s3.amazonaws.com/" @pytest.mark.requires_botocore -class S3TestCase(unittest.TestCase): +class TestS3: download_handler_cls: type = S3DownloadHandler # test use same example keys than amazon developer guide @@ -870,7 +867,7 @@ class S3TestCase(unittest.TestCase): AWS_ACCESS_KEY_ID = "0PN5J17HBGZHT7JJ3X82" AWS_SECRET_ACCESS_KEY = "uV3F3YluFJax1cknvbcGwgjvx4QpvB+leU8dUj2o" - def setUp(self): + def setup_method(self): crawler = get_crawler() s3reqh = build_from_crawler( S3DownloadHandler, @@ -897,17 +894,13 @@ class S3TestCase(unittest.TestCase): yield def test_extra_kw(self): - try: - crawler = get_crawler() + crawler = get_crawler() + with pytest.raises((TypeError, NotConfigured)): build_from_crawler( S3DownloadHandler, crawler, extra_kw=True, ) - except Exception as e: - self.assertIsInstance(e, (TypeError, NotConfigured)) - else: - raise AssertionError def test_request_signing1(self): # gets an object from the johnsmith bucket. @@ -915,9 +908,9 @@ class S3TestCase(unittest.TestCase): req = Request("s3://johnsmith/photos/puppy.jpg", headers={"Date": date}) with self._mocked_date(date): httpreq = self.download_request(req, self.spider) - self.assertEqual( - httpreq.headers["Authorization"], - b"AWS 0PN5J17HBGZHT7JJ3X82:xXjDGYUmKxnwqr5KXNPGldn5LbA=", + assert ( + httpreq.headers["Authorization"] + == b"AWS 0PN5J17HBGZHT7JJ3X82:xXjDGYUmKxnwqr5KXNPGldn5LbA=" ) def test_request_signing2(self): @@ -934,9 +927,9 @@ class S3TestCase(unittest.TestCase): ) with self._mocked_date(date): httpreq = self.download_request(req, self.spider) - self.assertEqual( - httpreq.headers["Authorization"], - b"AWS 0PN5J17HBGZHT7JJ3X82:hcicpDDvL9SsO6AkvxqmIWkmOuQ=", + assert ( + httpreq.headers["Authorization"] + == b"AWS 0PN5J17HBGZHT7JJ3X82:hcicpDDvL9SsO6AkvxqmIWkmOuQ=" ) def test_request_signing3(self): @@ -952,9 +945,9 @@ class S3TestCase(unittest.TestCase): ) with self._mocked_date(date): httpreq = self.download_request(req, self.spider) - self.assertEqual( - httpreq.headers["Authorization"], - b"AWS 0PN5J17HBGZHT7JJ3X82:jsRt/rhG+Vtp88HrYL706QhE4w4=", + assert ( + httpreq.headers["Authorization"] + == b"AWS 0PN5J17HBGZHT7JJ3X82:jsRt/rhG+Vtp88HrYL706QhE4w4=" ) def test_request_signing4(self): @@ -963,9 +956,9 @@ class S3TestCase(unittest.TestCase): req = Request("s3://johnsmith/?acl", method="GET", headers={"Date": date}) with self._mocked_date(date): httpreq = self.download_request(req, self.spider) - self.assertEqual( - httpreq.headers["Authorization"], - b"AWS 0PN5J17HBGZHT7JJ3X82:thdUi9VAkzhkniLj96JIrOPGi0g=", + assert ( + httpreq.headers["Authorization"] + == b"AWS 0PN5J17HBGZHT7JJ3X82:thdUi9VAkzhkniLj96JIrOPGi0g=" ) def test_request_signing6(self): @@ -991,9 +984,9 @@ class S3TestCase(unittest.TestCase): ) with self._mocked_date(date): httpreq = self.download_request(req, self.spider) - self.assertEqual( - httpreq.headers["Authorization"], - b"AWS 0PN5J17HBGZHT7JJ3X82:C0FlOtU8Ylb9KDTpZqYkZPX91iI=", + assert ( + httpreq.headers["Authorization"] + == b"AWS 0PN5J17HBGZHT7JJ3X82:C0FlOtU8Ylb9KDTpZqYkZPX91iI=" ) def test_request_signing7(self): @@ -1006,13 +999,13 @@ class S3TestCase(unittest.TestCase): ) with self._mocked_date(date): httpreq = self.download_request(req, self.spider) - self.assertEqual( - httpreq.headers["Authorization"], - b"AWS 0PN5J17HBGZHT7JJ3X82:+CfvG8EZ3YccOrRVMXNaK2eKZmM=", + assert ( + httpreq.headers["Authorization"] + == b"AWS 0PN5J17HBGZHT7JJ3X82:+CfvG8EZ3YccOrRVMXNaK2eKZmM=" ) -class BaseFTPTestCase(unittest.TestCase): +class TestFTPBase(unittest.TestCase): username = "scrapy" password = "passwd" req_meta = {"ftp_user": username, "ftp_password": password} @@ -1068,10 +1061,10 @@ class BaseFTPTestCase(unittest.TestCase): d = self.download_handler.download_request(request, None) def _test(r): - self.assertEqual(r.status, 200) - self.assertEqual(r.body, b"I have the power!") - self.assertEqual(r.headers, {b"Local Filename": [b""], b"Size": [b"17"]}) - self.assertIsNone(r.protocol) + 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 return self._add_test_callbacks(d, _test) @@ -1083,9 +1076,9 @@ class BaseFTPTestCase(unittest.TestCase): d = self.download_handler.download_request(request, None) def _test(r): - self.assertEqual(r.status, 200) - self.assertEqual(r.body, b"Moooooooooo power!") - self.assertEqual(r.headers, {b"Local Filename": [b""], b"Size": [b"18"]}) + assert r.status == 200 + assert r.body == b"Moooooooooo power!" + assert r.headers == {b"Local Filename": [b""], b"Size": [b"18"]} return self._add_test_callbacks(d, _test) @@ -1096,7 +1089,7 @@ class BaseFTPTestCase(unittest.TestCase): d = self.download_handler.download_request(request, None) def _test(r): - self.assertEqual(r.status, 404) + assert r.status == 404 return self._add_test_callbacks(d, _test) @@ -1111,12 +1104,10 @@ class BaseFTPTestCase(unittest.TestCase): d = self.download_handler.download_request(request, None) def _test(r): - self.assertEqual(r.body, fname_bytes) - self.assertEqual( - r.headers, {b"Local Filename": [fname_bytes], b"Size": [b"17"]} - ) - self.assertTrue(local_fname.exists()) - self.assertEqual(local_fname.read_bytes(), b"I have the power!") + 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() return self._add_test_callbacks(d, _test) @@ -1131,7 +1122,7 @@ class BaseFTPTestCase(unittest.TestCase): d = self.download_handler.download_request(request, None) def _test(r): - self.assertEqual(type(r), response_class) + assert type(r) is response_class # pylint: disable=unidiomatic-typecheck local_fname.unlink() return self._add_test_callbacks(d, _test) @@ -1143,7 +1134,7 @@ class BaseFTPTestCase(unittest.TestCase): return self._test_response_class("html-file-without-extension", HtmlResponse) -class FTPTestCase(BaseFTPTestCase): +class TestFTP(TestFTPBase): def test_invalid_credentials(self): if self.reactor_pytest == "asyncio" and sys.platform == "win32": raise unittest.SkipTest( @@ -1157,12 +1148,12 @@ class FTPTestCase(BaseFTPTestCase): d = self.download_handler.download_request(request, None) def _test(r): - self.assertEqual(r.type, ConnectionLost) + assert r.type == ConnectionLost return self._add_test_callbacks(d, errback=_test) -class AnonymousFTPTestCase(BaseFTPTestCase): +class TestAnonymousFTP(TestFTPBase): username = "anonymous" req_meta = {} @@ -1188,7 +1179,7 @@ class AnonymousFTPTestCase(BaseFTPTestCase): shutil.rmtree(self.directory) -class DataURITestCase(unittest.TestCase): +class TestDataURI(unittest.TestCase): def setUp(self): crawler = get_crawler() self.download_handler = build_from_crawler(DataURIDownloadHandler, crawler) @@ -1199,44 +1190,44 @@ class DataURITestCase(unittest.TestCase): uri = "data:,A%20brief%20note" def _test(response): - self.assertEqual(response.url, uri) - self.assertFalse(response.headers) + assert response.url == uri + assert not response.headers request = Request(uri) return self.download_request(request, self.spider).addCallback(_test) def test_default_mediatype_encoding(self): def _test(response): - self.assertEqual(response.text, "A brief note") - self.assertEqual(type(response), responsetypes.from_mimetype("text/plain")) - self.assertEqual(response.encoding, "US-ASCII") + assert response.text == "A brief note" + assert type(response) is responsetypes.from_mimetype("text/plain") # pylint: disable=unidiomatic-typecheck + assert response.encoding == "US-ASCII" request = Request("data:,A%20brief%20note") return self.download_request(request, self.spider).addCallback(_test) def test_default_mediatype(self): def _test(response): - self.assertEqual(response.text, "\u038e\u03a3\u038e") - self.assertEqual(type(response), responsetypes.from_mimetype("text/plain")) - self.assertEqual(response.encoding, "iso-8859-7") + 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" request = Request("data:;charset=iso-8859-7,%be%d3%be") return self.download_request(request, self.spider).addCallback(_test) def test_text_charset(self): def _test(response): - self.assertEqual(response.text, "\u038e\u03a3\u038e") - self.assertEqual(response.body, b"\xbe\xd3\xbe") - self.assertEqual(response.encoding, "iso-8859-7") + assert response.text == "\u038e\u03a3\u038e" + assert response.body == b"\xbe\xd3\xbe" + assert response.encoding == "iso-8859-7" request = Request("data:text/plain;charset=iso-8859-7,%be%d3%be") return self.download_request(request, self.spider).addCallback(_test) def test_mediatype_parameters(self): def _test(response): - self.assertEqual(response.text, "\u038e\u03a3\u038e") - self.assertEqual(type(response), responsetypes.from_mimetype("text/plain")) - self.assertEqual(response.encoding, "utf-8") + assert response.text == "\u038e\u03a3\u038e" + assert type(response) is responsetypes.from_mimetype("text/plain") # pylint: disable=unidiomatic-typecheck + assert response.encoding == "utf-8" request = Request( "data:text/plain;foo=%22foo;bar%5C%22%22;" @@ -1247,14 +1238,14 @@ class DataURITestCase(unittest.TestCase): def test_base64(self): def _test(response): - self.assertEqual(response.text, "Hello, world.") + assert response.text == "Hello, world." request = Request("data:text/plain;base64,SGVsbG8sIHdvcmxkLg%3D%3D") return self.download_request(request, self.spider).addCallback(_test) def test_protocol(self): def _test(response): - self.assertIsNone(response.protocol) + assert response.protocol is None request = Request("data:,") return self.download_request(request, self.spider).addCallback(_test) diff --git a/tests/test_downloader_handlers_http2.py b/tests/test_downloader_handlers_http2.py index 17d5c2d0a..c74c09cbb 100644 --- a/tests/test_downloader_handlers_http2.py +++ b/tests/test_downloader_handlers_http2.py @@ -4,7 +4,6 @@ from unittest import mock import pytest from testfixtures import LogCapture from twisted.internet import defer, error, reactor -from twisted.trial import unittest from twisted.web import server from twisted.web.error import SchemeNotSupported from twisted.web.http import H2_ENABLED @@ -28,25 +27,25 @@ class BaseTestClasses: # A hack to prevent tests from the imported classes to run here too. # See https://stackoverflow.com/q/1323455/113586 for other ways. from tests.test_downloader_handlers import ( - Http11MockServerTestCase as Http11MockServerTestCase, + TestHttp11MockServer as TestHttp11MockServer, ) from tests.test_downloader_handlers import ( - Http11ProxyTestCase as Http11ProxyTestCase, + TestHttp11Proxy as TestHttp11Proxy, ) from tests.test_downloader_handlers import ( - Https11CustomCiphers as Https11CustomCiphers, + TestHttps11 as TestHttps11, ) from tests.test_downloader_handlers import ( - Https11InvalidDNSId as Https11InvalidDNSId, + TestHttps11CustomCiphers as TestHttps11CustomCiphers, ) from tests.test_downloader_handlers import ( - Https11InvalidDNSPattern as Https11InvalidDNSPattern, + TestHttps11InvalidDNSId as TestHttps11InvalidDNSId, ) from tests.test_downloader_handlers import ( - Https11TestCase as Https11TestCase, + TestHttps11InvalidDNSPattern as TestHttps11InvalidDNSPattern, ) from tests.test_downloader_handlers import ( - Https11WrongHostnameTestCase as Https11WrongHostnameTestCase, + TestHttps11WrongHostname as TestHttps11WrongHostname, ) @@ -56,7 +55,7 @@ def _get_dh() -> type[DownloadHandlerProtocol]: return H2DownloadHandler -class Https2TestCase(BaseTestClasses.Https11TestCase): +class TestHttps2(BaseTestClasses.TestHttps11): scheme = "https" HTTP2_DATALOSS_SKIP_REASON = "Content-Length mismatch raises InvalidBodyLengthError" @@ -97,22 +96,22 @@ class Https2TestCase(BaseTestClasses.Https11TestCase): yield self.assertFailure(d, SchemeNotSupported) def test_download_broken_content_cause_data_loss(self, url="broken"): - raise unittest.SkipTest(self.HTTP2_DATALOSS_SKIP_REASON) + pytest.skip(self.HTTP2_DATALOSS_SKIP_REASON) def test_download_broken_chunked_content_cause_data_loss(self): - raise unittest.SkipTest(self.HTTP2_DATALOSS_SKIP_REASON) + pytest.skip(self.HTTP2_DATALOSS_SKIP_REASON) def test_download_broken_content_allow_data_loss(self, url="broken"): - raise unittest.SkipTest(self.HTTP2_DATALOSS_SKIP_REASON) + pytest.skip(self.HTTP2_DATALOSS_SKIP_REASON) def test_download_broken_chunked_content_allow_data_loss(self): - raise unittest.SkipTest(self.HTTP2_DATALOSS_SKIP_REASON) + pytest.skip(self.HTTP2_DATALOSS_SKIP_REASON) def test_download_broken_content_allow_data_loss_via_setting(self, url="broken"): - raise unittest.SkipTest(self.HTTP2_DATALOSS_SKIP_REASON) + pytest.skip(self.HTTP2_DATALOSS_SKIP_REASON) def test_download_broken_chunked_content_allow_data_loss_via_setting(self): - raise unittest.SkipTest(self.HTTP2_DATALOSS_SKIP_REASON) + pytest.skip(self.HTTP2_DATALOSS_SKIP_REASON) def test_concurrent_requests_same_domain(self): spider = Spider("foo") @@ -180,31 +179,31 @@ class Https2TestCase(BaseTestClasses.Https11TestCase): return d -class Https2WrongHostnameTestCase(BaseTestClasses.Https11WrongHostnameTestCase): +class Https2WrongHostnameTestCase(BaseTestClasses.TestHttps11WrongHostname): @property def download_handler_cls(self) -> type[DownloadHandlerProtocol]: return _get_dh() -class Https2InvalidDNSId(BaseTestClasses.Https11InvalidDNSId): +class Https2InvalidDNSId(BaseTestClasses.TestHttps11InvalidDNSId): @property def download_handler_cls(self) -> type[DownloadHandlerProtocol]: return _get_dh() -class Https2InvalidDNSPattern(BaseTestClasses.Https11InvalidDNSPattern): +class Https2InvalidDNSPattern(BaseTestClasses.TestHttps11InvalidDNSPattern): @property def download_handler_cls(self) -> type[DownloadHandlerProtocol]: return _get_dh() -class Https2CustomCiphers(BaseTestClasses.Https11CustomCiphers): +class Https2CustomCiphers(BaseTestClasses.TestHttps11CustomCiphers): @property def download_handler_cls(self) -> type[DownloadHandlerProtocol]: return _get_dh() -class Http2MockServerTestCase(BaseTestClasses.Http11MockServerTestCase): +class Http2MockServerTestCase(BaseTestClasses.TestHttp11MockServer): """HTTP 2.0 test case with MockServer""" settings_dict = { @@ -215,7 +214,7 @@ class Http2MockServerTestCase(BaseTestClasses.Http11MockServerTestCase): is_secure = True -class Https2ProxyTestCase(BaseTestClasses.Http11ProxyTestCase): +class Https2ProxyTestCase(BaseTestClasses.TestHttp11Proxy): # only used for HTTPS tests keyfile = "keys/localhost.key" certfile = "keys/localhost.crt" diff --git a/tests/test_exporters.py b/tests/test_exporters.py index 48728e078..f55cb6c97 100644 --- a/tests/test_exporters.py +++ b/tests/test_exporters.py @@ -54,11 +54,11 @@ class CustomFieldDataclass: age: int = dataclasses.field(metadata={"serializer": custom_serializer}) -class BaseItemExporterTest(unittest.TestCase): +class TestBaseItemExporter: item_class: type = MyItem custom_field_item_class: type = CustomFieldItem - def setUp(self): + def setup_method(self): self.i = self.item_class(name="John\xa3", age="22") self.output = BytesIO() self.ie = self._get_exporter() @@ -72,7 +72,7 @@ class BaseItemExporterTest(unittest.TestCase): def _assert_expected_item(self, exported_dict): for k, v in exported_dict.items(): exported_dict[k] = to_unicode(v) - self.assertEqual(self.i, self.item_class(**exported_dict)) + assert self.i == self.item_class(**exported_dict) def _get_nonstring_types_item(self): return { @@ -105,45 +105,40 @@ class BaseItemExporterTest(unittest.TestCase): def test_serialize_field(self): a = ItemAdapter(self.i) res = self.ie.serialize_field(a.get_field_meta("name"), "name", a["name"]) - self.assertEqual(res, "John\xa3") + assert res == "John\xa3" res = self.ie.serialize_field(a.get_field_meta("age"), "age", a["age"]) - self.assertEqual(res, "22") + assert res == "22" def test_fields_to_export(self): ie = self._get_exporter(fields_to_export=["name"]) - self.assertEqual( - list(ie._get_serialized_fields(self.i)), [("name", "John\xa3")] - ) + assert list(ie._get_serialized_fields(self.i)) == [("name", "John\xa3")] ie = self._get_exporter(fields_to_export=["name"], encoding="latin-1") _, name = next(iter(ie._get_serialized_fields(self.i))) assert isinstance(name, str) - self.assertEqual(name, "John\xa3") + assert name == "John\xa3" ie = self._get_exporter(fields_to_export={"name": "名稱"}) - self.assertEqual( - list(ie._get_serialized_fields(self.i)), [("名稱", "John\xa3")] - ) + assert list(ie._get_serialized_fields(self.i)) == [("名稱", "John\xa3")] def test_field_custom_serializer(self): i = self.custom_field_item_class(name="John\xa3", age="22") a = ItemAdapter(i) ie = self._get_exporter() - self.assertEqual( - ie.serialize_field(a.get_field_meta("name"), "name", a["name"]), "John\xa3" - ) - self.assertEqual( - ie.serialize_field(a.get_field_meta("age"), "age", a["age"]), "24" + assert ( + ie.serialize_field(a.get_field_meta("name"), "name", a["name"]) + == "John\xa3" ) + assert ie.serialize_field(a.get_field_meta("age"), "age", a["age"]) == "24" -class BaseItemExporterDataclassTest(BaseItemExporterTest): +class TestBaseItemExporterDataclass(TestBaseItemExporter): item_class = MyDataClass custom_field_item_class = CustomFieldDataclass -class PythonItemExporterTest(BaseItemExporterTest): +class TestPythonItemExporter(TestBaseItemExporter): def _get_exporter(self, **kwargs): return PythonItemExporter(**kwargs) @@ -157,16 +152,13 @@ class PythonItemExporterTest(BaseItemExporterTest): i3 = self.item_class(name="Jesus", age=i2) ie = self._get_exporter() exported = ie.export_item(i3) - self.assertEqual(type(exported), dict) - self.assertEqual( - exported, - { - "age": {"age": {"age": "22", "name": "Joseph"}, "name": "Maria"}, - "name": "Jesus", - }, - ) - self.assertEqual(type(exported["age"]), dict) - self.assertEqual(type(exported["age"]["age"]), dict) + assert isinstance(exported, dict) + assert exported == { + "age": {"age": {"age": "22", "name": "Joseph"}, "name": "Maria"}, + "name": "Jesus", + } + assert isinstance(exported["age"], dict) + assert isinstance(exported["age"]["age"], dict) def test_export_list(self): i1 = self.item_class(name="Joseph", age="22") @@ -174,15 +166,12 @@ class PythonItemExporterTest(BaseItemExporterTest): i3 = self.item_class(name="Jesus", age=[i2]) ie = self._get_exporter() exported = ie.export_item(i3) - self.assertEqual( - exported, - { - "age": [{"age": [{"age": "22", "name": "Joseph"}], "name": "Maria"}], - "name": "Jesus", - }, - ) - self.assertEqual(type(exported["age"][0]), dict) - self.assertEqual(type(exported["age"][0]["age"][0]), dict) + assert exported == { + "age": [{"age": [{"age": "22", "name": "Joseph"}], "name": "Maria"}], + "name": "Jesus", + } + assert isinstance(exported["age"][0], dict) + assert isinstance(exported["age"][0]["age"][0], dict) def test_export_item_dict_list(self): i1 = self.item_class(name="Joseph", age="22") @@ -190,29 +179,26 @@ class PythonItemExporterTest(BaseItemExporterTest): i3 = self.item_class(name="Jesus", age=[i2]) ie = self._get_exporter() exported = ie.export_item(i3) - self.assertEqual( - exported, - { - "age": [{"age": [{"age": "22", "name": "Joseph"}], "name": "Maria"}], - "name": "Jesus", - }, - ) - self.assertEqual(type(exported["age"][0]), dict) - self.assertEqual(type(exported["age"][0]["age"][0]), dict) + assert exported == { + "age": [{"age": [{"age": "22", "name": "Joseph"}], "name": "Maria"}], + "name": "Jesus", + } + assert isinstance(exported["age"][0], dict) + assert isinstance(exported["age"][0]["age"][0], dict) def test_nonstring_types_item(self): item = self._get_nonstring_types_item() ie = self._get_exporter() exported = ie.export_item(item) - self.assertEqual(exported, item) + assert exported == item -class PythonItemExporterDataclassTest(PythonItemExporterTest): +class TestPythonItemExporterDataclass(TestPythonItemExporter): item_class = MyDataClass custom_field_item_class = CustomFieldDataclass -class PprintItemExporterTest(BaseItemExporterTest): +class TestPprintItemExporter(TestBaseItemExporter): def _get_exporter(self, **kwargs): return PprintItemExporter(self.output, **kwargs) @@ -222,12 +208,12 @@ class PprintItemExporterTest(BaseItemExporterTest): ) -class PprintItemExporterDataclassTest(PprintItemExporterTest): +class TestPprintItemExporterDataclass(TestPprintItemExporter): item_class = MyDataClass custom_field_item_class = CustomFieldDataclass -class PickleItemExporterTest(BaseItemExporterTest): +class TestPickleItemExporter(TestBaseItemExporter): def _get_exporter(self, **kwargs): return PickleItemExporter(self.output, **kwargs) @@ -245,8 +231,8 @@ class PickleItemExporterTest(BaseItemExporterTest): ie.finish_exporting() del ie # See the first “del self.ie” in this file for context. f.seek(0) - self.assertEqual(self.item_class(**pickle.load(f)), i1) - self.assertEqual(self.item_class(**pickle.load(f)), i2) + assert self.item_class(**pickle.load(f)) == i1 + assert self.item_class(**pickle.load(f)) == i2 def test_nonstring_types_item(self): item = self._get_nonstring_types_item() @@ -256,15 +242,15 @@ class PickleItemExporterTest(BaseItemExporterTest): ie.export_item(item) ie.finish_exporting() del ie # See the first “del self.ie” in this file for context. - self.assertEqual(pickle.loads(fp.getvalue()), item) + assert pickle.loads(fp.getvalue()) == item -class PickleItemExporterDataclassTest(PickleItemExporterTest): +class TestPickleItemExporterDataclass(TestPickleItemExporter): item_class = MyDataClass custom_field_item_class = CustomFieldDataclass -class MarshalItemExporterTest(BaseItemExporterTest): +class TestMarshalItemExporter(TestBaseItemExporter): def _get_exporter(self, **kwargs): self.output = tempfile.TemporaryFile() return MarshalItemExporter(self.output, **kwargs) @@ -283,15 +269,15 @@ class MarshalItemExporterTest(BaseItemExporterTest): ie.finish_exporting() del ie # See the first “del self.ie” in this file for context. fp.seek(0) - self.assertEqual(marshal.load(fp), item) + assert marshal.load(fp) == item -class MarshalItemExporterDataclassTest(MarshalItemExporterTest): +class TestMarshalItemExporterDataclass(TestMarshalItemExporter): item_class = MyDataClass custom_field_item_class = CustomFieldDataclass -class CsvItemExporterTest(BaseItemExporterTest): +class TestCsvItemExporter(TestBaseItemExporter): def _get_exporter(self, **kwargs): self.output = tempfile.TemporaryFile() return CsvItemExporter(self.output, **kwargs) @@ -303,7 +289,7 @@ class CsvItemExporterTest(BaseItemExporterTest): for line in to_unicode(csv).splitlines(True) ] - return self.assertEqual(split_csv(first), split_csv(second), msg=msg) + assert split_csv(first) == split_csv(second), msg def _check_output(self): self.output.seek(0) @@ -406,12 +392,12 @@ class CsvItemExporterTest(BaseItemExporterTest): ) -class CsvItemExporterDataclassTest(CsvItemExporterTest): +class TestCsvItemExporterDataclass(TestCsvItemExporter): item_class = MyDataClass custom_field_item_class = CustomFieldDataclass -class XmlItemExporterTest(BaseItemExporterTest): +class TestXmlItemExporter(TestBaseItemExporter): def _get_exporter(self, **kwargs): return XmlItemExporter(self.output, **kwargs) @@ -426,7 +412,7 @@ class XmlItemExporterTest(BaseItemExporterTest): doc = lxml.etree.fromstring(xmlcontent) return xmltuple(doc) - return self.assertEqual(xmlsplit(first), xmlsplit(second), msg) + assert xmlsplit(first) == xmlsplit(second), msg def assertExportResult(self, item, expected_value): fp = BytesIO() @@ -517,12 +503,12 @@ class XmlItemExporterTest(BaseItemExporterTest): ) -class XmlItemExporterDataclassTest(XmlItemExporterTest): +class TestXmlItemExporterDataclass(TestXmlItemExporter): item_class = MyDataClass custom_field_item_class = CustomFieldDataclass -class JsonLinesItemExporterTest(BaseItemExporterTest): +class TestJsonLinesItemExporter(TestBaseItemExporter): _expected_nested: Any = { "name": "Jesus", "age": {"name": "Maria", "age": {"name": "Joseph", "age": "22"}}, @@ -533,7 +519,7 @@ class JsonLinesItemExporterTest(BaseItemExporterTest): def _check_output(self): exported = json.loads(to_unicode(self.output.getvalue().strip())) - self.assertEqual(exported, ItemAdapter(self.i).asdict()) + assert exported == ItemAdapter(self.i).asdict() def test_nested_item(self): i1 = self.item_class(name="Joseph", age="22") @@ -544,7 +530,7 @@ class JsonLinesItemExporterTest(BaseItemExporterTest): self.ie.finish_exporting() del self.ie # See the first “del self.ie” in this file for context. exported = json.loads(to_unicode(self.output.getvalue())) - self.assertEqual(exported, self._expected_nested) + assert exported == self._expected_nested def test_extra_keywords(self): self.ie = self._get_exporter(sort_keys=True) @@ -561,23 +547,23 @@ class JsonLinesItemExporterTest(BaseItemExporterTest): del self.ie # See the first “del self.ie” in this file for context. exported = json.loads(to_unicode(self.output.getvalue())) item["time"] = str(item["time"]) - self.assertEqual(exported, item) + assert exported == item -class JsonLinesItemExporterDataclassTest(JsonLinesItemExporterTest): +class TestJsonLinesItemExporterDataclass(TestJsonLinesItemExporter): item_class = MyDataClass custom_field_item_class = CustomFieldDataclass -class JsonItemExporterTest(JsonLinesItemExporterTest): - _expected_nested = [JsonLinesItemExporterTest._expected_nested] +class TestJsonItemExporter(TestJsonLinesItemExporter): + _expected_nested = [TestJsonLinesItemExporter._expected_nested] def _get_exporter(self, **kwargs): return JsonItemExporter(self.output, **kwargs) def _check_output(self): exported = json.loads(to_unicode(self.output.getvalue().strip())) - self.assertEqual(exported, [ItemAdapter(self.i).asdict()]) + assert exported == [ItemAdapter(self.i).asdict()] def assertTwoItemsExported(self, item): self.ie.start_exporting() @@ -586,9 +572,7 @@ class JsonItemExporterTest(JsonLinesItemExporterTest): self.ie.finish_exporting() del self.ie # See the first “del self.ie” in this file for context. exported = json.loads(to_unicode(self.output.getvalue())) - self.assertEqual( - exported, [ItemAdapter(item).asdict(), ItemAdapter(item).asdict()] - ) + assert exported == [ItemAdapter(item).asdict(), ItemAdapter(item).asdict()] def test_two_items(self): self.assertTwoItemsExported(self.i) @@ -609,7 +593,7 @@ class JsonItemExporterTest(JsonLinesItemExporterTest): self.ie.export_item(i3) self.ie.finish_exporting() exported = json.loads(to_unicode(self.output.getvalue())) - self.assertEqual(exported, [dict(i1), dict(i3)]) + assert exported == [dict(i1), dict(i3)] def test_nested_item(self): i1 = self.item_class(name="Joseph\xa3", age="22") @@ -624,7 +608,7 @@ class JsonItemExporterTest(JsonLinesItemExporterTest): "name": "Jesus", "age": {"name": "Maria", "age": ItemAdapter(i1).asdict()}, } - self.assertEqual(exported, [expected]) + assert exported == [expected] def test_nested_dict_item(self): i1 = {"name": "Joseph\xa3", "age": "22"} @@ -636,7 +620,7 @@ class JsonItemExporterTest(JsonLinesItemExporterTest): del self.ie # See the first “del self.ie” in this file for context. exported = json.loads(to_unicode(self.output.getvalue())) expected = {"name": "Jesus", "age": {"name": "Maria", "age": i1}} - self.assertEqual(exported, [expected]) + assert exported == [expected] def test_nonstring_types_item(self): item = self._get_nonstring_types_item() @@ -646,10 +630,10 @@ class JsonItemExporterTest(JsonLinesItemExporterTest): del self.ie # See the first “del self.ie” in this file for context. exported = json.loads(to_unicode(self.output.getvalue())) item["time"] = str(item["time"]) - self.assertEqual(exported, [item]) + assert exported == [item] -class JsonItemExporterToBytesTest(BaseItemExporterTest): +class TestJsonItemExporterToBytes(TestBaseItemExporter): def _get_exporter(self, **kwargs): kwargs["encoding"] = "latin" return JsonItemExporter(self.output, **kwargs) @@ -665,18 +649,18 @@ class JsonItemExporterToBytesTest(BaseItemExporterTest): self.ie.export_item(i3) self.ie.finish_exporting() exported = json.loads(to_unicode(self.output.getvalue(), encoding="latin")) - self.assertEqual(exported, [dict(i1), dict(i3)]) + assert exported == [dict(i1), dict(i3)] -class JsonItemExporterDataclassTest(JsonItemExporterTest): +class TestJsonItemExporterDataclass(TestJsonItemExporter): item_class = MyDataClass custom_field_item_class = CustomFieldDataclass -class CustomExporterItemTest(unittest.TestCase): +class TestCustomExporterItem: item_class: type = MyItem - def setUp(self): + def setup_method(self): if self.item_class is None: raise unittest.SkipTest("item class is None") @@ -691,17 +675,13 @@ class CustomExporterItemTest(unittest.TestCase): a = ItemAdapter(i) ie = CustomItemExporter() - self.assertEqual( - ie.serialize_field(a.get_field_meta("name"), "name", a["name"]), "John" - ) - self.assertEqual( - ie.serialize_field(a.get_field_meta("age"), "age", a["age"]), "23" - ) + assert ie.serialize_field(a.get_field_meta("name"), "name", a["name"]) == "John" + assert ie.serialize_field(a.get_field_meta("age"), "age", a["age"]) == "23" i2 = {"name": "John", "age": "22"} - self.assertEqual(ie.serialize_field({}, "name", i2["name"]), "John") - self.assertEqual(ie.serialize_field({}, "age", i2["age"]), "23") + assert ie.serialize_field({}, "name", i2["name"]) == "John" + assert ie.serialize_field({}, "age", i2["age"]) == "23" -class CustomExporterDataclassTest(CustomExporterItemTest): +class TestCustomExporterDataclass(TestCustomExporterItem): item_class = MyDataClass diff --git a/tests/test_feedexport.py b/tests/test_feedexport.py index 8e008ab98..44cd10ec3 100644 --- a/tests/test_feedexport.py +++ b/tests/test_feedexport.py @@ -88,7 +88,7 @@ def mock_google_cloud_storage() -> tuple[Any, Any, Any]: return (client_mock, bucket_mock, blob_mock) -class FileFeedStorageTest(unittest.TestCase): +class TestFileFeedStorage(unittest.TestCase): def test_store_file_uri(self): path = Path(self.mktemp()).resolve() uri = path_to_file_uri(str(path)) @@ -137,14 +137,14 @@ class FileFeedStorageTest(unittest.TestCase): file = storage.open(spider) file.write(b"content") yield storage.store(file) - self.assertTrue(path.exists()) + assert path.exists() try: - self.assertEqual(path.read_bytes(), expected_content) + assert path.read_bytes() == expected_content finally: path.unlink() -class FTPFeedStorageTest(unittest.TestCase): +class TestFTPFeedStorage(unittest.TestCase): def get_test_spider(self, settings=None): class TestSpider(scrapy.Spider): name = "test_spider" @@ -166,9 +166,9 @@ class FTPFeedStorageTest(unittest.TestCase): return storage.store(file) def _assert_stored(self, path: Path, content): - self.assertTrue(path.exists()) + assert path.exists() try: - self.assertEqual(path.read_bytes(), content) + assert path.read_bytes() == content finally: path.unlink() @@ -216,10 +216,10 @@ class FTPFeedStorageTest(unittest.TestCase): # RFC3986: 3.2.1. User Information pw_quoted = quote(string.punctuation, safe="") st = FTPFeedStorage(f"ftp://foo:{pw_quoted}@example.com/some_path", {}) - self.assertEqual(st.password, string.punctuation) + assert st.password == string.punctuation -class BlockingFeedStorageTest(unittest.TestCase): +class TestBlockingFeedStorage: def get_test_spider(self, settings=None): class TestSpider(scrapy.Spider): name = "test_spider" @@ -232,7 +232,7 @@ class BlockingFeedStorageTest(unittest.TestCase): tmp = b.open(self.get_test_spider()) tmp_path = Path(tmp.name).parent - self.assertEqual(str(tmp_path), tempfile.gettempdir()) + assert str(tmp_path) == tempfile.gettempdir() def test_temp_file(self): b = BlockingFeedStorage() @@ -241,7 +241,7 @@ class BlockingFeedStorageTest(unittest.TestCase): spider = self.get_test_spider({"FEED_TEMPDIR": str(tests_path)}) tmp = b.open(spider) tmp_path = Path(tmp.name).parent - self.assertEqual(tmp_path, tests_path) + assert tmp_path == tests_path def test_invalid_folder(self): b = BlockingFeedStorage() @@ -255,7 +255,7 @@ class BlockingFeedStorageTest(unittest.TestCase): @pytest.mark.requires_boto3 -class S3FeedStorageTest(unittest.TestCase): +class TestS3FeedStorage(unittest.TestCase): def test_parse_credentials(self): aws_credentials = { "AWS_ACCESS_KEY_ID": "settings_key", @@ -268,9 +268,9 @@ class S3FeedStorageTest(unittest.TestCase): crawler, "s3://mybucket/export.csv", ) - self.assertEqual(storage.access_key, "settings_key") - self.assertEqual(storage.secret_key, "settings_secret") - self.assertEqual(storage.session_token, "settings_token") + assert storage.access_key == "settings_key" + assert storage.secret_key == "settings_secret" + assert storage.session_token == "settings_token" # Instantiate directly storage = S3FeedStorage( "s3://mybucket/export.csv", @@ -278,17 +278,17 @@ class S3FeedStorageTest(unittest.TestCase): aws_credentials["AWS_SECRET_ACCESS_KEY"], session_token=aws_credentials["AWS_SESSION_TOKEN"], ) - self.assertEqual(storage.access_key, "settings_key") - self.assertEqual(storage.secret_key, "settings_secret") - self.assertEqual(storage.session_token, "settings_token") + assert storage.access_key == "settings_key" + assert storage.secret_key == "settings_secret" + assert storage.session_token == "settings_token" # URI priority > settings priority storage = S3FeedStorage( "s3://uri_key:uri_secret@mybucket/export.csv", aws_credentials["AWS_ACCESS_KEY_ID"], aws_credentials["AWS_SECRET_ACCESS_KEY"], ) - self.assertEqual(storage.access_key, "uri_key") - self.assertEqual(storage.secret_key, "uri_secret") + assert storage.access_key == "uri_key" + assert storage.secret_key == "uri_secret" @defer.inlineCallbacks def test_store(self): @@ -306,24 +306,23 @@ class S3FeedStorageTest(unittest.TestCase): storage.s3_client = mock.MagicMock() yield storage.store(file) - self.assertEqual( - storage.s3_client.upload_fileobj.call_args, - mock.call(Bucket=bucket, Key=key, Fileobj=file), + assert storage.s3_client.upload_fileobj.call_args == mock.call( + Bucket=bucket, Key=key, Fileobj=file ) def test_init_without_acl(self): storage = S3FeedStorage("s3://mybucket/export.csv", "access_key", "secret_key") - self.assertEqual(storage.access_key, "access_key") - self.assertEqual(storage.secret_key, "secret_key") - self.assertEqual(storage.acl, None) + assert storage.access_key == "access_key" + assert storage.secret_key == "secret_key" + assert storage.acl is None def test_init_with_acl(self): storage = S3FeedStorage( "s3://mybucket/export.csv", "access_key", "secret_key", "custom-acl" ) - self.assertEqual(storage.access_key, "access_key") - self.assertEqual(storage.secret_key, "secret_key") - self.assertEqual(storage.acl, "custom-acl") + assert storage.access_key == "access_key" + assert storage.secret_key == "secret_key" + assert storage.acl == "custom-acl" def test_init_with_endpoint_url(self): storage = S3FeedStorage( @@ -332,9 +331,9 @@ class S3FeedStorageTest(unittest.TestCase): "secret_key", endpoint_url="https://example.com", ) - self.assertEqual(storage.access_key, "access_key") - self.assertEqual(storage.secret_key, "secret_key") - self.assertEqual(storage.endpoint_url, "https://example.com") + assert storage.access_key == "access_key" + assert storage.secret_key == "secret_key" + assert storage.endpoint_url == "https://example.com" def test_init_with_region_name(self): region_name = "ap-east-1" @@ -344,10 +343,10 @@ class S3FeedStorageTest(unittest.TestCase): "secret_key", region_name=region_name, ) - self.assertEqual(storage.access_key, "access_key") - self.assertEqual(storage.secret_key, "secret_key") - self.assertEqual(storage.region_name, region_name) - self.assertEqual(storage.s3_client._client_config.region_name, region_name) + assert storage.access_key == "access_key" + assert storage.secret_key == "secret_key" + assert storage.region_name == region_name + assert storage.s3_client._client_config.region_name == region_name def test_from_crawler_without_acl(self): settings = { @@ -359,9 +358,9 @@ class S3FeedStorageTest(unittest.TestCase): crawler, "s3://mybucket/export.csv", ) - self.assertEqual(storage.access_key, "access_key") - self.assertEqual(storage.secret_key, "secret_key") - self.assertEqual(storage.acl, None) + assert storage.access_key == "access_key" + assert storage.secret_key == "secret_key" + assert storage.acl is None def test_without_endpoint_url(self): settings = { @@ -373,9 +372,9 @@ class S3FeedStorageTest(unittest.TestCase): crawler, "s3://mybucket/export.csv", ) - self.assertEqual(storage.access_key, "access_key") - self.assertEqual(storage.secret_key, "secret_key") - self.assertEqual(storage.endpoint_url, None) + assert storage.access_key == "access_key" + assert storage.secret_key == "secret_key" + assert storage.endpoint_url is None def test_without_region_name(self): settings = { @@ -387,9 +386,9 @@ class S3FeedStorageTest(unittest.TestCase): crawler, "s3://mybucket/export.csv", ) - self.assertEqual(storage.access_key, "access_key") - self.assertEqual(storage.secret_key, "secret_key") - self.assertEqual(storage.s3_client._client_config.region_name, "us-east-1") + assert storage.access_key == "access_key" + assert storage.secret_key == "secret_key" + assert storage.s3_client._client_config.region_name == "us-east-1" def test_from_crawler_with_acl(self): settings = { @@ -402,9 +401,9 @@ class S3FeedStorageTest(unittest.TestCase): crawler, "s3://mybucket/export.csv", ) - self.assertEqual(storage.access_key, "access_key") - self.assertEqual(storage.secret_key, "secret_key") - self.assertEqual(storage.acl, "custom-acl") + assert storage.access_key == "access_key" + assert storage.secret_key == "secret_key" + assert storage.acl == "custom-acl" def test_from_crawler_with_endpoint_url(self): settings = { @@ -414,9 +413,9 @@ class S3FeedStorageTest(unittest.TestCase): } crawler = get_crawler(settings_dict=settings) storage = S3FeedStorage.from_crawler(crawler, "s3://mybucket/export.csv") - self.assertEqual(storage.access_key, "access_key") - self.assertEqual(storage.secret_key, "secret_key") - self.assertEqual(storage.endpoint_url, "https://example.com") + assert storage.access_key == "access_key" + assert storage.secret_key == "secret_key" + assert storage.endpoint_url == "https://example.com" def test_from_crawler_with_region_name(self): region_name = "ap-east-1" @@ -427,10 +426,10 @@ class S3FeedStorageTest(unittest.TestCase): } crawler = get_crawler(settings_dict=settings) storage = S3FeedStorage.from_crawler(crawler, "s3://mybucket/export.csv") - self.assertEqual(storage.access_key, "access_key") - self.assertEqual(storage.secret_key, "secret_key") - self.assertEqual(storage.region_name, region_name) - self.assertEqual(storage.s3_client._client_config.region_name, region_name) + assert storage.access_key == "access_key" + assert storage.secret_key == "secret_key" + assert storage.region_name == region_name + assert storage.s3_client._client_config.region_name == region_name @defer.inlineCallbacks def test_store_without_acl(self): @@ -439,9 +438,9 @@ class S3FeedStorageTest(unittest.TestCase): "access_key", "secret_key", ) - self.assertEqual(storage.access_key, "access_key") - self.assertEqual(storage.secret_key, "secret_key") - self.assertEqual(storage.acl, None) + assert storage.access_key == "access_key" + assert storage.secret_key == "secret_key" + assert storage.acl is None storage.s3_client = mock.MagicMock() yield storage.store(BytesIO(b"test file")) @@ -450,28 +449,28 @@ class S3FeedStorageTest(unittest.TestCase): .get("ExtraArgs", {}) .get("ACL") ) - self.assertIsNone(acl) + assert acl is None @defer.inlineCallbacks def test_store_with_acl(self): storage = S3FeedStorage( "s3://mybucket/export.csv", "access_key", "secret_key", "custom-acl" ) - self.assertEqual(storage.access_key, "access_key") - self.assertEqual(storage.secret_key, "secret_key") - self.assertEqual(storage.acl, "custom-acl") + assert storage.access_key == "access_key" + assert storage.secret_key == "secret_key" + assert storage.acl == "custom-acl" storage.s3_client = mock.MagicMock() yield storage.store(BytesIO(b"test file")) acl = storage.s3_client.upload_fileobj.call_args[1]["ExtraArgs"]["ACL"] - self.assertEqual(acl, "custom-acl") + assert acl == "custom-acl" def test_overwrite_default(self): with LogCapture() as log: S3FeedStorage( "s3://mybucket/export.csv", "access_key", "secret_key", "custom-acl" ) - self.assertNotIn("S3 does not support appending to files", str(log)) + assert "S3 does not support appending to files" not in str(log) def test_overwrite_false(self): with LogCapture() as log: @@ -482,10 +481,10 @@ class S3FeedStorageTest(unittest.TestCase): "custom-acl", feed_options={"overwrite": False}, ) - self.assertIn("S3 does not support appending to files", str(log)) + assert "S3 does not support appending to files" in str(log) -class GCSFeedStorageTest(unittest.TestCase): +class TestGCSFeedStorage(unittest.TestCase): def test_parse_settings(self): try: from google.cloud.storage import Client # noqa: F401 @@ -543,7 +542,7 @@ class GCSFeedStorageTest(unittest.TestCase): def test_overwrite_default(self): with LogCapture() as log: GCSFeedStorage("gs://mybucket/export.csv", "myproject-123", "custom-acl") - self.assertNotIn("GCS does not support appending to files", str(log)) + assert "GCS does not support appending to files" not in str(log) def test_overwrite_false(self): with LogCapture() as log: @@ -553,10 +552,10 @@ class GCSFeedStorageTest(unittest.TestCase): "custom-acl", feed_options={"overwrite": False}, ) - self.assertIn("GCS does not support appending to files", str(log)) + assert "GCS does not support appending to files" in str(log) -class StdoutFeedStorageTest(unittest.TestCase): +class TestStdoutFeedStorage(unittest.TestCase): @defer.inlineCallbacks def test_store(self): out = BytesIO() @@ -564,20 +563,21 @@ class StdoutFeedStorageTest(unittest.TestCase): file = storage.open(scrapy.Spider("default")) file.write(b"content") yield storage.store(file) - self.assertEqual(out.getvalue(), b"content") + assert out.getvalue() == b"content" def test_overwrite_default(self): with LogCapture() as log: StdoutFeedStorage("stdout:") - self.assertNotIn( - "Standard output (stdout) storage does not support overwriting", str(log) + assert ( + "Standard output (stdout) storage does not support overwriting" + not in str(log) ) def test_overwrite_true(self): with LogCapture() as log: StdoutFeedStorage("stdout:", feed_options={"overwrite": True}) - self.assertIn( - "Standard output (stdout) storage does not support overwriting", str(log) + assert "Standard output (stdout) storage does not support overwriting" in str( + log ) @@ -639,7 +639,7 @@ class LogOnStoreFileStorage: file.close() -class FeedExportTestBase(ABC, unittest.TestCase): +class TestFeedExportBase(ABC, unittest.TestCase): class MyItem(scrapy.Item): foo = scrapy.Field() egg = scrapy.Field() @@ -769,7 +769,7 @@ class ExceptionJsonItemExporter(JsonItemExporter): raise RuntimeError("foo") -class FeedExportTest(FeedExportTestBase): +class TestFeedExport(TestFeedExportBase): @defer.inlineCallbacks def run_and_export(self, spider_cls, settings): """Run spider with specified settings; return exported data.""" @@ -812,8 +812,8 @@ class FeedExportTest(FeedExportTestBase): ) data = yield self.exported_data(items, settings) reader = csv.DictReader(to_unicode(data["csv"]).splitlines()) - self.assertEqual(reader.fieldnames, list(header)) - self.assertEqual(rows, list(reader)) + assert reader.fieldnames == list(header) + assert rows == list(reader) @defer.inlineCallbacks def assertExportedJsonLines(self, items, rows, settings=None): @@ -828,7 +828,7 @@ class FeedExportTest(FeedExportTestBase): data = yield self.exported_data(items, settings) parsed = [json.loads(to_unicode(line)) for line in data["jl"].splitlines()] rows = [{k: v for k, v in row.items() if v} for row in rows] - self.assertEqual(rows, parsed) + assert rows == parsed @defer.inlineCallbacks def assertExportedXml(self, items, rows, settings=None): @@ -844,7 +844,7 @@ class FeedExportTest(FeedExportTestBase): rows = [{k: v for k, v in row.items() if v} for row in rows] root = lxml.etree.fromstring(data["xml"]) got_rows = [{e.tag: e.text for e in it} for it in root.findall("item")] - self.assertEqual(rows, got_rows) + assert rows == got_rows @defer.inlineCallbacks def assertExportedMultiple(self, items, rows, settings=None): @@ -862,10 +862,10 @@ class FeedExportTest(FeedExportTestBase): # XML root = lxml.etree.fromstring(data["xml"]) xml_rows = [{e.tag: e.text for e in it} for it in root.findall("item")] - self.assertEqual(rows, xml_rows) + assert rows == xml_rows # JSON json_rows = json.loads(to_unicode(data["json"])) - self.assertEqual(rows, json_rows) + assert rows == json_rows @defer.inlineCallbacks def assertExportedPickle(self, items, rows, settings=None): @@ -882,7 +882,7 @@ class FeedExportTest(FeedExportTestBase): import pickle result = self._load_until_eof(data["pickle"], load_func=pickle.load) - self.assertEqual(expected, result) + assert result == expected @defer.inlineCallbacks def assertExportedMarshal(self, items, rows, settings=None): @@ -899,7 +899,7 @@ class FeedExportTest(FeedExportTestBase): import marshal result = self._load_until_eof(data["marshal"], load_func=marshal.load) - self.assertEqual(expected, result) + assert result == expected @defer.inlineCallbacks def test_stats_file_success(self): @@ -912,12 +912,8 @@ class FeedExportTest(FeedExportTestBase): } crawler = get_crawler(ItemSpider, settings) yield crawler.crawl(mockserver=self.mockserver) - self.assertIn( - "feedexport/success_count/FileFeedStorage", crawler.stats.get_stats() - ) - self.assertEqual( - crawler.stats.get_value("feedexport/success_count/FileFeedStorage"), 1 - ) + assert "feedexport/success_count/FileFeedStorage" in crawler.stats.get_stats() + assert crawler.stats.get_value("feedexport/success_count/FileFeedStorage") == 1 @defer.inlineCallbacks def test_stats_file_failed(self): @@ -934,12 +930,8 @@ class FeedExportTest(FeedExportTestBase): side_effect=KeyError("foo"), ): yield crawler.crawl(mockserver=self.mockserver) - self.assertIn( - "feedexport/failed_count/FileFeedStorage", crawler.stats.get_stats() - ) - self.assertEqual( - crawler.stats.get_value("feedexport/failed_count/FileFeedStorage"), 1 - ) + assert "feedexport/failed_count/FileFeedStorage" in crawler.stats.get_stats() + assert crawler.stats.get_value("feedexport/failed_count/FileFeedStorage") == 1 @defer.inlineCallbacks def test_stats_multiple_file(self): @@ -956,17 +948,11 @@ class FeedExportTest(FeedExportTestBase): crawler = get_crawler(ItemSpider, settings) with mock.patch.object(S3FeedStorage, "store"): yield crawler.crawl(mockserver=self.mockserver) - self.assertIn( - "feedexport/success_count/FileFeedStorage", crawler.stats.get_stats() - ) - self.assertIn( - "feedexport/success_count/StdoutFeedStorage", crawler.stats.get_stats() - ) - self.assertEqual( - crawler.stats.get_value("feedexport/success_count/FileFeedStorage"), 1 - ) - self.assertEqual( - crawler.stats.get_value("feedexport/success_count/StdoutFeedStorage"), 1 + assert "feedexport/success_count/FileFeedStorage" in crawler.stats.get_stats() + assert "feedexport/success_count/StdoutFeedStorage" in crawler.stats.get_stats() + assert crawler.stats.get_value("feedexport/success_count/FileFeedStorage") == 1 + assert ( + crawler.stats.get_value("feedexport/success_count/StdoutFeedStorage") == 1 ) @defer.inlineCallbacks @@ -993,7 +979,7 @@ class FeedExportTest(FeedExportTestBase): "FEED_STORE_EMPTY": False, } data = yield self.exported_no_data(settings) - self.assertEqual(None, data[fmt]) + assert data[fmt] is None @defer.inlineCallbacks def test_start_finish_exporting_items(self): @@ -1012,8 +998,8 @@ class FeedExportTest(FeedExportTestBase): with mock.patch("scrapy.extensions.feedexport.FeedSlot", InstrumentedFeedSlot): _ = yield self.exported_data(items, settings) - self.assertFalse(listener.start_without_finish) - self.assertFalse(listener.finish_without_start) + assert not listener.start_without_finish + assert not listener.finish_without_start @defer.inlineCallbacks def test_start_finish_exporting_no_items(self): @@ -1030,8 +1016,8 @@ class FeedExportTest(FeedExportTestBase): with mock.patch("scrapy.extensions.feedexport.FeedSlot", InstrumentedFeedSlot): _ = yield self.exported_data(items, settings) - self.assertFalse(listener.start_without_finish) - self.assertFalse(listener.finish_without_start) + assert not listener.start_without_finish + assert not listener.finish_without_start @defer.inlineCallbacks def test_start_finish_exporting_items_exception(self): @@ -1051,8 +1037,8 @@ class FeedExportTest(FeedExportTestBase): with mock.patch("scrapy.extensions.feedexport.FeedSlot", InstrumentedFeedSlot): _ = yield self.exported_data(items, settings) - self.assertFalse(listener.start_without_finish) - self.assertFalse(listener.finish_without_start) + assert not listener.start_without_finish + assert not listener.finish_without_start @defer.inlineCallbacks def test_start_finish_exporting_no_items_exception(self): @@ -1070,8 +1056,8 @@ class FeedExportTest(FeedExportTestBase): with mock.patch("scrapy.extensions.feedexport.FeedSlot", InstrumentedFeedSlot): _ = yield self.exported_data(items, settings) - self.assertFalse(listener.start_without_finish) - self.assertFalse(listener.finish_without_start) + assert not listener.start_without_finish + assert not listener.finish_without_start @defer.inlineCallbacks def test_export_no_items_store_empty(self): @@ -1091,7 +1077,7 @@ class FeedExportTest(FeedExportTestBase): "FEED_EXPORT_INDENT": None, } data = yield self.exported_no_data(settings) - self.assertEqual(expctd, data[fmt]) + assert expctd == data[fmt] @defer.inlineCallbacks def test_export_no_items_multiple_feeds(self): @@ -1109,7 +1095,7 @@ class FeedExportTest(FeedExportTestBase): with LogCapture() as log: yield self.exported_no_data(settings) - self.assertEqual(str(log).count("Storage.store is called"), 0) + assert str(log).count("Storage.store is called") == 0 @defer.inlineCallbacks def test_export_multiple_item_classes(self): @@ -1238,7 +1224,7 @@ class FeedExportTest(FeedExportTestBase): data = yield self.exported_data(items, settings) for fmt, expected in formats.items(): - self.assertEqual(expected, data[fmt]) + assert data[fmt] == expected @defer.inlineCallbacks def test_export_based_on_custom_filters(self): @@ -1297,7 +1283,7 @@ class FeedExportTest(FeedExportTestBase): data = yield self.exported_data(items, settings) for fmt, expected in formats.items(): - self.assertEqual(expected, data[fmt]) + assert data[fmt] == expected @defer.inlineCallbacks def test_export_dicts(self): @@ -1371,7 +1357,7 @@ class FeedExportTest(FeedExportTestBase): "FEED_EXPORT_INDENT": None, } data = yield self.exported_data(items, settings) - self.assertEqual(expected, data[fmt]) + assert data[fmt] == expected formats = { "json": b'[{"foo": "Test\xd6"}]', @@ -1392,7 +1378,7 @@ class FeedExportTest(FeedExportTestBase): "FEED_EXPORT_ENCODING": "latin-1", } data = yield self.exported_data(items, settings) - self.assertEqual(expected, data[fmt]) + assert data[fmt] == expected @defer.inlineCallbacks def test_export_multiple_configs(self): @@ -1432,7 +1418,7 @@ class FeedExportTest(FeedExportTestBase): data = yield self.exported_data(items, settings) for fmt, expected in formats.items(): - self.assertEqual(expected, data[fmt]) + assert data[fmt] == expected @defer.inlineCallbacks def test_export_indentation(self): @@ -1588,7 +1574,7 @@ class FeedExportTest(FeedExportTestBase): }, } data = yield self.exported_data(items, settings) - self.assertEqual(row["expected"], data[row["format"]]) + assert data[row["format"]] == row["expected"] @defer.inlineCallbacks def test_init_exporters_storages_with_crawler(self): @@ -1600,8 +1586,8 @@ class FeedExportTest(FeedExportTestBase): }, } yield self.exported_data(items=[], settings=settings) - self.assertTrue(FromCrawlerCsvItemExporter.init_with_crawler) - self.assertTrue(FromCrawlerFileFeedStorage.init_with_crawler) + assert FromCrawlerCsvItemExporter.init_with_crawler + assert FromCrawlerFileFeedStorage.init_with_crawler @defer.inlineCallbacks def test_str_uri(self): @@ -1610,7 +1596,7 @@ class FeedExportTest(FeedExportTestBase): "FEEDS": {str(self._random_temp_filename()): {"format": "csv"}}, } data = yield self.exported_no_data(settings) - self.assertEqual(data["csv"], b"") + assert data["csv"] == b"" @defer.inlineCallbacks def test_multiple_feeds_success_logs_blocking_feed_storage(self): @@ -1631,7 +1617,7 @@ class FeedExportTest(FeedExportTestBase): print(log) for fmt in ["json", "xml", "csv"]: - self.assertIn(f"Stored {fmt} feed (2 items)", str(log)) + assert f"Stored {fmt} feed (2 items)" in str(log) @defer.inlineCallbacks def test_multiple_feeds_failing_logs_blocking_feed_storage(self): @@ -1652,7 +1638,7 @@ class FeedExportTest(FeedExportTestBase): print(log) for fmt in ["json", "xml", "csv"]: - self.assertIn(f"Error storing {fmt} feed (2 items)", str(log)) + assert f"Error storing {fmt} feed (2 items)" in str(log) @defer.inlineCallbacks def test_extend_kwargs(self): @@ -1689,7 +1675,7 @@ class FeedExportTest(FeedExportTestBase): } data = yield self.exported_data(items, settings) - self.assertEqual(row["expected"], data[feed_options["format"]]) + assert data[feed_options["format"]] == row["expected"] @defer.inlineCallbacks def test_storage_file_no_postprocessing(self): @@ -1711,7 +1697,7 @@ class FeedExportTest(FeedExportTestBase): "FEED_STORAGES": {"file": Storage}, } yield self.exported_no_data(settings) - self.assertIs(Storage.open_file, Storage.store_file) + assert Storage.open_file is Storage.store_file @defer.inlineCallbacks def test_storage_file_postprocessing(self): @@ -1741,11 +1727,11 @@ class FeedExportTest(FeedExportTestBase): "FEED_STORAGES": {"file": Storage}, } yield self.exported_no_data(settings) - self.assertIs(Storage.open_file, Storage.store_file) - self.assertFalse(Storage.file_was_closed) + assert Storage.open_file is Storage.store_file + assert not Storage.file_was_closed -class FeedPostProcessedExportsTest(FeedExportTestBase): +class TestFeedPostProcessedExports(TestFeedExportBase): items = [{"foo": "bar"}] expected = b"foo\r\nbar\r\n" @@ -1827,7 +1813,7 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): try: gzip.decompress(data[filename]) except OSError: - self.fail("Received invalid gzip data.") + pytest.fail("Received invalid gzip data.") @defer.inlineCallbacks def test_gzip_plugin_compresslevel(self): @@ -1863,8 +1849,8 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): for filename, compressed in filename_to_compressed.items(): result = gzip.decompress(data[filename]) - self.assertEqual(compressed, data[filename]) - self.assertEqual(self.expected, result) + assert compressed == data[filename] + assert result == self.expected @defer.inlineCallbacks def test_gzip_plugin_mtime(self): @@ -1898,8 +1884,8 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): for filename, compressed in filename_to_compressed.items(): result = gzip.decompress(data[filename]) - self.assertEqual(compressed, data[filename]) - self.assertEqual(self.expected, result) + assert compressed == data[filename] + assert result == self.expected @defer.inlineCallbacks def test_gzip_plugin_filename(self): @@ -1933,8 +1919,8 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): for filename, compressed in filename_to_compressed.items(): result = gzip.decompress(data[filename]) - self.assertEqual(compressed, data[filename]) - self.assertEqual(self.expected, result) + assert compressed == data[filename] + assert result == self.expected @defer.inlineCallbacks def test_lzma_plugin(self): @@ -1953,7 +1939,7 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): try: lzma.decompress(data[filename]) except lzma.LZMAError: - self.fail("Received invalid lzma data.") + pytest.fail("Received invalid lzma data.") @defer.inlineCallbacks def test_lzma_plugin_format(self): @@ -1985,8 +1971,8 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): for filename, compressed in filename_to_compressed.items(): result = lzma.decompress(data[filename]) - self.assertEqual(compressed, data[filename]) - self.assertEqual(self.expected, result) + assert compressed == data[filename] + assert result == self.expected @defer.inlineCallbacks def test_lzma_plugin_check(self): @@ -2018,8 +2004,8 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): for filename, compressed in filename_to_compressed.items(): result = lzma.decompress(data[filename]) - self.assertEqual(compressed, data[filename]) - self.assertEqual(self.expected, result) + assert compressed == data[filename] + assert result == self.expected @defer.inlineCallbacks def test_lzma_plugin_preset(self): @@ -2051,8 +2037,8 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): for filename, compressed in filename_to_compressed.items(): result = lzma.decompress(data[filename]) - self.assertEqual(compressed, data[filename]) - self.assertEqual(self.expected, result) + assert compressed == data[filename] + assert result == self.expected @defer.inlineCallbacks def test_lzma_plugin_filters(self): @@ -2075,9 +2061,9 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): } data = yield self.exported_data(self.items, settings) - self.assertEqual(compressed, data[filename]) + assert compressed == data[filename] result = lzma.decompress(data[filename]) - self.assertEqual(self.expected, result) + assert result == self.expected @defer.inlineCallbacks def test_bz2_plugin(self): @@ -2096,7 +2082,7 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): try: bz2.decompress(data[filename]) except OSError: - self.fail("Received invalid bz2 data.") + pytest.fail("Received invalid bz2 data.") @defer.inlineCallbacks def test_bz2_plugin_compresslevel(self): @@ -2128,8 +2114,8 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): for filename, compressed in filename_to_compressed.items(): result = bz2.decompress(data[filename]) - self.assertEqual(compressed, data[filename]) - self.assertEqual(self.expected, result) + assert compressed == data[filename] + assert result == self.expected @defer.inlineCallbacks def test_custom_plugin(self): @@ -2145,7 +2131,7 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): } data = yield self.exported_data(self.items, settings) - self.assertEqual(self.expected, data[filename]) + assert data[filename] == self.expected @defer.inlineCallbacks def test_custom_plugin_with_parameter(self): @@ -2163,7 +2149,7 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): } data = yield self.exported_data(self.items, settings) - self.assertEqual(expected, data[filename]) + assert data[filename] == expected @defer.inlineCallbacks def test_custom_plugin_with_compression(self): @@ -2208,7 +2194,7 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): for filename, decompressor in filename_to_decompressor.items(): result = decompressor(data[filename]) - self.assertEqual(expected, result) + assert result == expected @defer.inlineCallbacks def test_exports_compatibility_with_postproc(self): @@ -2262,10 +2248,10 @@ class FeedPostProcessedExportsTest(FeedExportTestBase): expected, result = self.items[0], marshal.loads(result) else: expected = filename_to_expected[filename] - self.assertEqual(expected, result) + assert result == expected -class BatchDeliveriesTest(FeedExportTestBase): +class TestBatchDeliveries(TestFeedExportBase): _file_mark = "_%(batch_time)s_#%(batch_id)02d_" @defer.inlineCallbacks @@ -2310,7 +2296,7 @@ class BatchDeliveriesTest(FeedExportTestBase): json.loads(to_unicode(batch_item)) for batch_item in batch.splitlines() ] expected_batch, rows = rows[:batch_size], rows[batch_size:] - self.assertEqual(expected_batch, got_batch) + assert got_batch == expected_batch @defer.inlineCallbacks def assertExportedCsv(self, items, header, rows, settings=None): @@ -2328,9 +2314,9 @@ class BatchDeliveriesTest(FeedExportTestBase): data = yield self.exported_data(items, settings) for batch in data["csv"]: got_batch = csv.DictReader(to_unicode(batch).splitlines()) - self.assertEqual(list(header), got_batch.fieldnames) + assert list(header) == got_batch.fieldnames expected_batch, rows = rows[:batch_size], rows[batch_size:] - self.assertEqual(expected_batch, list(got_batch)) + assert list(got_batch) == expected_batch @defer.inlineCallbacks def assertExportedXml(self, items, rows, settings=None): @@ -2351,7 +2337,7 @@ class BatchDeliveriesTest(FeedExportTestBase): root = lxml.etree.fromstring(batch) got_batch = [{e.tag: e.text for e in it} for it in root.findall("item")] expected_batch, rows = rows[:batch_size], rows[batch_size:] - self.assertEqual(expected_batch, got_batch) + assert got_batch == expected_batch @defer.inlineCallbacks def assertExportedMultiple(self, items, rows, settings=None): @@ -2377,13 +2363,13 @@ class BatchDeliveriesTest(FeedExportTestBase): root = lxml.etree.fromstring(batch) got_batch = [{e.tag: e.text for e in it} for it in root.findall("item")] expected_batch, xml_rows = xml_rows[:batch_size], xml_rows[batch_size:] - self.assertEqual(expected_batch, got_batch) + assert got_batch == expected_batch # JSON json_rows = rows.copy() for batch in data["json"]: got_batch = json.loads(batch.decode("utf-8")) expected_batch, json_rows = json_rows[:batch_size], json_rows[batch_size:] - self.assertEqual(expected_batch, got_batch) + assert got_batch == expected_batch @defer.inlineCallbacks def assertExportedPickle(self, items, rows, settings=None): @@ -2405,7 +2391,7 @@ class BatchDeliveriesTest(FeedExportTestBase): for batch in data["pickle"]: got_batch = self._load_until_eof(batch, load_func=pickle.load) expected_batch, rows = rows[:batch_size], rows[batch_size:] - self.assertEqual(expected_batch, got_batch) + assert got_batch == expected_batch @defer.inlineCallbacks def assertExportedMarshal(self, items, rows, settings=None): @@ -2427,7 +2413,7 @@ class BatchDeliveriesTest(FeedExportTestBase): for batch in data["marshal"]: got_batch = self._load_until_eof(batch, load_func=marshal.load) expected_batch, rows = rows[:batch_size], rows[batch_size:] - self.assertEqual(expected_batch, got_batch) + assert got_batch == expected_batch @defer.inlineCallbacks def test_export_items(self): @@ -2472,7 +2458,7 @@ class BatchDeliveriesTest(FeedExportTestBase): } data = yield self.exported_no_data(settings) data = dict(data) - self.assertEqual(0, len(data[fmt])) + assert len(data[fmt]) == 0 @defer.inlineCallbacks def test_export_no_items_store_empty(self): @@ -2496,7 +2482,7 @@ class BatchDeliveriesTest(FeedExportTestBase): } data = yield self.exported_no_data(settings) data = dict(data) - self.assertEqual(expctd, data[fmt][0]) + assert data[fmt][0] == expctd @defer.inlineCallbacks def test_export_multiple_configs(self): @@ -2552,7 +2538,7 @@ class BatchDeliveriesTest(FeedExportTestBase): data = yield self.exported_data(items, settings) for fmt, expected in formats.items(): for expected_batch, got_batch in zip(expected, data[fmt]): - self.assertEqual(expected_batch, got_batch) + assert got_batch == expected_batch @defer.inlineCallbacks def test_batch_item_count_feeds_setting(self): @@ -2576,7 +2562,7 @@ class BatchDeliveriesTest(FeedExportTestBase): data = yield self.exported_data(items, settings) for fmt, expected in formats.items(): for expected_batch, got_batch in zip(expected, data[fmt]): - self.assertEqual(expected_batch, got_batch) + assert got_batch == expected_batch @defer.inlineCallbacks def test_batch_path_differ(self): @@ -2598,7 +2584,7 @@ class BatchDeliveriesTest(FeedExportTestBase): "FEED_EXPORT_BATCH_ITEM_COUNT": 1, } data = yield self.exported_data(items, settings) - self.assertEqual(len(items), len(data["json"])) + assert len(items) == len(data["json"]) @defer.inlineCallbacks def test_stats_batch_file_success(self): @@ -2614,12 +2600,8 @@ class BatchDeliveriesTest(FeedExportTestBase): } crawler = get_crawler(ItemSpider, settings) yield crawler.crawl(total=2, mockserver=self.mockserver) - self.assertIn( - "feedexport/success_count/FileFeedStorage", crawler.stats.get_stats() - ) - self.assertEqual( - crawler.stats.get_value("feedexport/success_count/FileFeedStorage"), 12 - ) + assert "feedexport/success_count/FileFeedStorage" in crawler.stats.get_stats() + assert crawler.stats.get_value("feedexport/success_count/FileFeedStorage") == 12 @pytest.mark.requires_boto3 @defer.inlineCallbacks @@ -2687,13 +2669,13 @@ class BatchDeliveriesTest(FeedExportTestBase): crawler = get_crawler(TestSpider, settings) yield crawler.crawl() - self.assertEqual(len(CustomS3FeedStorage.stubs), len(items)) + assert len(CustomS3FeedStorage.stubs) == len(items) for stub in CustomS3FeedStorage.stubs[:-1]: stub.assert_no_pending_responses() # Test that the FeedExporer sends the feed_exporter_closed and feed_slot_closed signals -class FeedExporterSignalsTest(unittest.TestCase): +class TestFeedExporterSignals: items = [ {"foo": "bar1", "egg": "spam1"}, {"foo": "bar2", "egg": "spam2", "baz": "quux2"}, @@ -2754,8 +2736,8 @@ class FeedExporterSignalsTest(unittest.TestCase): self.feed_exporter_closed_signal_handler, self.feed_slot_closed_signal_handler, ) - self.assertTrue(self.feed_slot_closed_received) - self.assertTrue(self.feed_exporter_closed_received) + assert self.feed_slot_closed_received + assert self.feed_exporter_closed_received def test_feed_exporter_signals_sent_deferred(self): self.feed_exporter_closed_received = False @@ -2765,11 +2747,11 @@ class FeedExporterSignalsTest(unittest.TestCase): self.feed_exporter_closed_signal_handler_deferred, self.feed_slot_closed_signal_handler_deferred, ) - self.assertTrue(self.feed_slot_closed_received) - self.assertTrue(self.feed_exporter_closed_received) + assert self.feed_slot_closed_received + assert self.feed_exporter_closed_received -class FeedExportInitTest(unittest.TestCase): +class TestFeedExportInit: def test_unsupported_storage(self): settings = { "FEEDS": { @@ -2803,7 +2785,7 @@ class FeedExportInitTest(unittest.TestCase): } crawler = get_crawler(settings_dict=settings) exporter = FeedExporter.from_crawler(crawler) - self.assertIsInstance(exporter, FeedExporter) + assert isinstance(exporter, FeedExporter) def test_relative_pathlib_as_uri(self): settings = { @@ -2815,13 +2797,14 @@ class FeedExportInitTest(unittest.TestCase): } crawler = get_crawler(settings_dict=settings) exporter = FeedExporter.from_crawler(crawler) - self.assertIsInstance(exporter, FeedExporter) + assert isinstance(exporter, FeedExporter) -class URIParamsTest: +class TestURIParams(ABC): spider_name = "uri_params_spider" deprecated_options = False + @abstractmethod def build_settings(self, uri="file:///tmp/foobar", uri_params=None): raise NotImplementedError @@ -2850,7 +2833,7 @@ class URIParamsTest: warnings.simplefilter("error", ScrapyDeprecationWarning) feed_exporter.open_spider(spider) - self.assertEqual(feed_exporter.slots[0].uri, f"file:///tmp/{self.spider_name}") + assert feed_exporter.slots[0].uri == f"file:///tmp/{self.spider_name}" def test_none(self): def uri_params(params, spider): @@ -2866,7 +2849,7 @@ class URIParamsTest: feed_exporter.open_spider(spider) - self.assertEqual(feed_exporter.slots[0].uri, f"file:///tmp/{self.spider_name}") + assert feed_exporter.slots[0].uri == f"file:///tmp/{self.spider_name}" def test_empty_dict(self): def uri_params(params, spider): @@ -2900,7 +2883,7 @@ class URIParamsTest: warnings.simplefilter("error", ScrapyDeprecationWarning) feed_exporter.open_spider(spider) - self.assertEqual(feed_exporter.slots[0].uri, f"file:///tmp/{self.spider_name}") + assert feed_exporter.slots[0].uri == f"file:///tmp/{self.spider_name}" def test_custom_param(self): def uri_params(params, spider): @@ -2917,10 +2900,10 @@ class URIParamsTest: warnings.simplefilter("error", ScrapyDeprecationWarning) feed_exporter.open_spider(spider) - self.assertEqual(feed_exporter.slots[0].uri, f"file:///tmp/{self.spider_name}") + assert feed_exporter.slots[0].uri == f"file:///tmp/{self.spider_name}" -class URIParamsSettingTest(URIParamsTest, unittest.TestCase): +class TestURIParamsSetting(TestURIParams): deprecated_options = True def build_settings(self, uri="file:///tmp/foobar", uri_params=None): @@ -2933,7 +2916,7 @@ class URIParamsSettingTest(URIParamsTest, unittest.TestCase): } -class URIParamsFeedOptionTest(URIParamsTest, unittest.TestCase): +class TestURIParamsFeedOption(TestURIParams): deprecated_options = False def build_settings(self, uri="file:///tmp/foobar", uri_params=None): diff --git a/tests/test_http2_client_protocol.py b/tests/test_http2_client_protocol.py index 0881bbeca..7c1b38877 100644 --- a/tests/test_http2_client_protocol.py +++ b/tests/test_http2_client_protocol.py @@ -185,7 +185,7 @@ def get_client_certificate( @skipIf(not H2_ENABLED, "HTTP/2 support in Twisted is not enabled") -class Https2ClientProtocolTestCase(TestCase): +class TestHttps2ClientProtocol(TestCase): scheme = "https" key_file = Path(__file__).parent / "keys" / "localhost.key" certificate_file = Path(__file__).parent / "keys" / "localhost.crt" @@ -277,14 +277,14 @@ class Https2ClientProtocolTestCase(TestCase): def _check_GET(self, request: Request, expected_body, expected_status): def check_response(response: Response): - self.assertEqual(response.status, expected_status) - self.assertEqual(response.body, expected_body) - self.assertEqual(response.request, request) + assert response.status == expected_status + assert response.body == expected_body + assert response.request == request content_length_header = response.headers.get("Content-Length") assert content_length_header is not None content_length = int(content_length_header) - self.assertEqual(len(response.body), content_length) + assert len(response.body) == content_length d = self.make_request(request) d.addCallback(check_response) @@ -325,35 +325,35 @@ class Https2ClientProtocolTestCase(TestCase): d = self.make_request(request) def assert_response(response: Response): - self.assertEqual(response.status, expected_status) - self.assertEqual(response.request, request) + assert response.status == expected_status + assert response.request == request content_length_header = response.headers.get("Content-Length") assert content_length_header is not None content_length = int(content_length_header) - self.assertEqual(len(response.body), content_length) + assert len(response.body) == content_length # Parse the body content_encoding_header = response.headers[b"Content-Encoding"] assert content_encoding_header is not None content_encoding = str(content_encoding_header, "utf-8") body = json.loads(str(response.body, content_encoding)) - self.assertIn("request-body", body) - self.assertIn("extra-data", body) - self.assertIn("request-headers", body) + assert "request-body" in body + assert "extra-data" in body + assert "request-headers" in body request_body = body["request-body"] - self.assertEqual(request_body, expected_request_body) + assert request_body == expected_request_body extra_data = body["extra-data"] - self.assertEqual(extra_data, expected_extra_data) + assert extra_data == expected_extra_data # Check if headers were sent successfully request_headers = body["request-headers"] for k, v in request.headers.items(): k_str = str(k, "utf-8") - self.assertIn(k_str, request_headers) - self.assertEqual(request_headers[k_str], str(v[0], "utf-8")) + assert k_str in request_headers + assert request_headers[k_str] == str(v[0], "utf-8") d.addCallback(assert_response) d.addErrback(self.fail) @@ -414,8 +414,8 @@ class Https2ClientProtocolTestCase(TestCase): request = Request(url=self.get_url("/get-data-html-large")) def assert_response(response: Response): - self.assertEqual(response.status, 499) - self.assertEqual(response.request, request) + assert response.status == 499 + assert response.request == request d = self.make_request(request) d.addCallback(assert_response) @@ -430,12 +430,12 @@ class Https2ClientProtocolTestCase(TestCase): ) def assert_cancelled_error(failure): - self.assertIsInstance(failure.value, CancelledError) + assert isinstance(failure.value, CancelledError) error_pattern = re.compile( rf"Cancelling download of {request.url}: received response " rf"size \(\d*\) larger than download max size \(1000\)" ) - self.assertEqual(len(re.findall(error_pattern, str(failure.value))), 1) + assert len(re.findall(error_pattern, str(failure.value))) == 1 d = self.make_request(request) d.addCallback(self.fail) @@ -448,14 +448,12 @@ class Https2ClientProtocolTestCase(TestCase): request = Request(url=self.get_url("/dataloss")) def assert_failure(failure: Failure): - self.assertTrue(len(failure.value.reasons) > 0) + assert len(failure.value.reasons) > 0 from h2.exceptions import InvalidBodyLengthError - self.assertTrue( - any( - isinstance(error, InvalidBodyLengthError) - for error in failure.value.reasons - ) + assert any( + isinstance(error, InvalidBodyLengthError) + for error in failure.value.reasons ) d = self.make_request(request) @@ -467,10 +465,10 @@ class Https2ClientProtocolTestCase(TestCase): request = Request(url=self.get_url("/no-content-length-header")) def assert_content_length(response: Response): - self.assertEqual(response.status, 200) - self.assertEqual(response.body, Data.NO_CONTENT_LENGTH) - self.assertEqual(response.request, request) - self.assertNotIn("Content-Length", response.headers) + assert response.status == 200 + assert response.body == Data.NO_CONTENT_LENGTH + assert response.request == request + assert "Content-Length" not in response.headers d = self.make_request(request) d.addCallback(assert_content_length) @@ -481,14 +479,12 @@ class Https2ClientProtocolTestCase(TestCase): def _check_log_warnsize(self, request, warn_pattern, expected_body): with self.assertLogs("scrapy.core.http2.stream", level="WARNING") as cm: response = yield self.make_request(request) - self.assertEqual(response.status, 200) - self.assertEqual(response.request, request) - self.assertEqual(response.body, expected_body) + assert response.status == 200 + assert response.request == request + assert response.body == expected_body # Check the warning is raised only once for this request - self.assertEqual( - sum(len(re.findall(warn_pattern, log)) for log in cm.output), 1 - ) + assert sum(len(re.findall(warn_pattern, log)) for log in cm.output) == 1 @inlineCallbacks def test_log_expected_warnsize(self): @@ -534,11 +530,11 @@ class Https2ClientProtocolTestCase(TestCase): d_list = [] def assert_inactive_stream(failure): - self.assertIsNotNone(failure.check(ResponseFailed)) + assert failure.check(ResponseFailed) is not None from scrapy.core.http2.stream import InactiveStreamClosed - self.assertTrue( - any(isinstance(e, InactiveStreamClosed) for e in failure.value.reasons) + assert any( + isinstance(e, InactiveStreamClosed) for e in failure.value.reasons ) # Send 100 request (we do not check the result) @@ -578,7 +574,7 @@ class Https2ClientProtocolTestCase(TestCase): assert content_encoding_header is not None content_encoding = str(content_encoding_header, "utf-8") data = json.loads(str(response.body, content_encoding)) - self.assertEqual(data, params) + assert data == params d = self.make_request(request) d.addCallback(assert_query_params) @@ -588,7 +584,7 @@ class Https2ClientProtocolTestCase(TestCase): def test_status_codes(self): def assert_response_status(response: Response, expected_status: int): - self.assertEqual(response.status, expected_status) + assert response.status == expected_status d_list = [] for status in [200, 404]: @@ -604,21 +600,18 @@ class Https2ClientProtocolTestCase(TestCase): request = Request(self.get_url("/status?n=200")) def assert_metadata(response: Response): - self.assertEqual(response.request, request) - self.assertIsInstance(response.certificate, Certificate) - assert response.certificate # typing - self.assertIsNotNone(response.certificate.original) - self.assertEqual( - response.certificate.getIssuer(), self.client_certificate.getIssuer() + assert response.request == request + assert isinstance(response.certificate, Certificate) + assert response.certificate.original is not None + assert ( + response.certificate.getIssuer() == self.client_certificate.getIssuer() ) - self.assertTrue( - response.certificate.getPublicKey().matches( - self.client_certificate.getPublicKey() - ) + assert response.certificate.getPublicKey().matches( + self.client_certificate.getPublicKey() ) - self.assertIsInstance(response.ip_address, IPv4Address) - self.assertEqual(str(response.ip_address), "127.0.0.1") + assert isinstance(response.ip_address, IPv4Address) + assert str(response.ip_address) == "127.0.0.1" d = self.make_request(request) d.addCallback(assert_metadata) @@ -632,11 +625,11 @@ class Https2ClientProtocolTestCase(TestCase): def assert_invalid_hostname(failure: Failure): from scrapy.core.http2.stream import InvalidHostname - self.assertIsNotNone(failure.check(InvalidHostname)) + assert failure.check(InvalidHostname) is not None error_msg = str(failure.value) - self.assertIn("localhost", error_msg) - self.assertIn("127.0.0.1", error_msg) - self.assertIn(str(request), error_msg) + assert "localhost" in error_msg + assert "127.0.0.1" in error_msg + assert str(request) in error_msg d = self.make_request(request) d.addCallback(self.fail) @@ -672,13 +665,13 @@ class Https2ClientProtocolTestCase(TestCase): from scrapy.core.http2.protocol import H2ClientProtocol if isinstance(err, TimeoutError): - self.assertIn( - f"Connection was IDLE for more than {H2ClientProtocol.IDLE_TIMEOUT}s", - str(err), + assert ( + f"Connection was IDLE for more than {H2ClientProtocol.IDLE_TIMEOUT}s" + in str(err) ) break else: - self.fail() + pytest.fail("No TimeoutError raised.") d.addCallback(self.fail) d.addErrback(assert_timeout_error) @@ -692,15 +685,15 @@ class Https2ClientProtocolTestCase(TestCase): d = self.make_request(request) def assert_request_headers(response: Response): - self.assertEqual(response.status, 200) - self.assertEqual(response.request, request) + assert response.status == 200 + assert response.request == request response_headers = json.loads(str(response.body, "utf-8")) - self.assertIsInstance(response_headers, dict) + assert isinstance(response_headers, dict) for k, v in request.headers.items(): k, v = str(k, "utf-8"), str(v[0], "utf-8") - self.assertIn(k, response_headers) - self.assertEqual(v, response_headers[k]) + assert k in response_headers + assert v == response_headers[k] d.addErrback(self.fail) d.addCallback(assert_request_headers) diff --git a/tests/test_http_cookies.py b/tests/test_http_cookies.py index 932644320..660b76d08 100644 --- a/tests/test_http_cookies.py +++ b/tests/test_http_cookies.py @@ -1,74 +1,72 @@ -from unittest import TestCase - from scrapy.http import Request, Response from scrapy.http.cookies import WrappedRequest, WrappedResponse from scrapy.utils.httpobj import urlparse_cached -class WrappedRequestTest(TestCase): - def setUp(self): +class TestWrappedRequest: + def setup_method(self): self.request = Request( "http://www.example.com/page.html", headers={"Content-Type": "text/html"} ) self.wrapped = WrappedRequest(self.request) def test_get_full_url(self): - self.assertEqual(self.wrapped.get_full_url(), self.request.url) - self.assertEqual(self.wrapped.full_url, self.request.url) + assert self.wrapped.get_full_url() == self.request.url + assert self.wrapped.full_url == self.request.url def test_get_host(self): - self.assertEqual(self.wrapped.get_host(), urlparse_cached(self.request).netloc) - self.assertEqual(self.wrapped.host, urlparse_cached(self.request).netloc) + assert self.wrapped.get_host() == urlparse_cached(self.request).netloc + assert self.wrapped.host == urlparse_cached(self.request).netloc def test_get_type(self): - self.assertEqual(self.wrapped.get_type(), urlparse_cached(self.request).scheme) - self.assertEqual(self.wrapped.type, urlparse_cached(self.request).scheme) + assert self.wrapped.get_type() == urlparse_cached(self.request).scheme + assert self.wrapped.type == urlparse_cached(self.request).scheme def test_is_unverifiable(self): - self.assertFalse(self.wrapped.is_unverifiable()) - self.assertFalse(self.wrapped.unverifiable) + assert not self.wrapped.is_unverifiable() + assert not self.wrapped.unverifiable def test_is_unverifiable2(self): self.request.meta["is_unverifiable"] = True - self.assertTrue(self.wrapped.is_unverifiable()) - self.assertTrue(self.wrapped.unverifiable) + assert self.wrapped.is_unverifiable() + assert self.wrapped.unverifiable def test_get_origin_req_host(self): - self.assertEqual(self.wrapped.origin_req_host, "www.example.com") + assert self.wrapped.origin_req_host == "www.example.com" def test_has_header(self): - self.assertTrue(self.wrapped.has_header("content-type")) - self.assertFalse(self.wrapped.has_header("xxxxx")) + assert self.wrapped.has_header("content-type") + assert not self.wrapped.has_header("xxxxx") def test_get_header(self): - self.assertEqual(self.wrapped.get_header("content-type"), "text/html") - self.assertEqual(self.wrapped.get_header("xxxxx", "def"), "def") - self.assertEqual(self.wrapped.get_header("xxxxx"), None) + assert self.wrapped.get_header("content-type") == "text/html" + assert self.wrapped.get_header("xxxxx", "def") == "def" + assert self.wrapped.get_header("xxxxx") is None wrapped = WrappedRequest( Request( "http://www.example.com/page.html", headers={"empty-binary-header": b""} ) ) - self.assertEqual(wrapped.get_header("empty-binary-header"), "") + assert wrapped.get_header("empty-binary-header") == "" def test_header_items(self): - self.assertEqual(self.wrapped.header_items(), [("Content-Type", ["text/html"])]) + assert self.wrapped.header_items() == [("Content-Type", ["text/html"])] def test_add_unredirected_header(self): self.wrapped.add_unredirected_header("hello", "world") - self.assertEqual(self.request.headers["hello"], b"world") + assert self.request.headers["hello"] == b"world" -class WrappedResponseTest(TestCase): - def setUp(self): +class TestWrappedResponse: + def setup_method(self): self.response = Response( "http://www.example.com/page.html", headers={"Content-TYpe": "text/html"} ) self.wrapped = WrappedResponse(self.response) def test_info(self): - self.assertIs(self.wrapped.info(), self.wrapped) + assert self.wrapped.info() is self.wrapped def test_get_all(self): # get_all result must be native string - self.assertEqual(self.wrapped.get_all("content-type"), ["text/html"]) + assert self.wrapped.get_all("content-type") == ["text/html"] diff --git a/tests/test_http_headers.py b/tests/test_http_headers.py index 0bbbcda46..2fcf9e83c 100644 --- a/tests/test_http_headers.py +++ b/tests/test_http_headers.py @@ -1,14 +1,13 @@ import copy -import unittest import pytest from scrapy.http import Headers -class HeadersTest(unittest.TestCase): +class TestHeaders: def assertSortedEqual(self, first, second, msg=None): - return self.assertEqual(sorted(first), sorted(second), msg) + assert sorted(first) == sorted(second), msg def test_basics(self): h = Headers({"Content-Type": "text/html", "Content-Length": 1234}) @@ -17,53 +16,53 @@ class HeadersTest(unittest.TestCase): with pytest.raises(KeyError): h["Accept"] - self.assertEqual(h.get("Accept"), None) - self.assertEqual(h.getlist("Accept"), []) + assert h.get("Accept") is None + assert h.getlist("Accept") == [] - self.assertEqual(h.get("Accept", "*/*"), b"*/*") - self.assertEqual(h.getlist("Accept", "*/*"), [b"*/*"]) - self.assertEqual( - h.getlist("Accept", ["text/html", "images/jpeg"]), - [b"text/html", b"images/jpeg"], - ) + assert h.get("Accept", "*/*") == b"*/*" + assert h.getlist("Accept", "*/*") == [b"*/*"] + assert h.getlist("Accept", ["text/html", "images/jpeg"]) == [ + b"text/html", + b"images/jpeg", + ] def test_single_value(self): h = Headers() h["Content-Type"] = "text/html" - self.assertEqual(h["Content-Type"], b"text/html") - self.assertEqual(h.get("Content-Type"), b"text/html") - self.assertEqual(h.getlist("Content-Type"), [b"text/html"]) + assert h["Content-Type"] == b"text/html" + assert h.get("Content-Type") == b"text/html" + assert h.getlist("Content-Type") == [b"text/html"] def test_multivalue(self): h = Headers() h["X-Forwarded-For"] = hlist = ["ip1", "ip2"] - self.assertEqual(h["X-Forwarded-For"], b"ip2") - self.assertEqual(h.get("X-Forwarded-For"), b"ip2") - self.assertEqual(h.getlist("X-Forwarded-For"), [b"ip1", b"ip2"]) + assert h["X-Forwarded-For"] == b"ip2" + assert h.get("X-Forwarded-For") == b"ip2" + assert h.getlist("X-Forwarded-For") == [b"ip1", b"ip2"] assert h.getlist("X-Forwarded-For") is not hlist def test_multivalue_for_one_header(self): h = Headers((("a", "b"), ("a", "c"))) - self.assertEqual(h["a"], b"c") - self.assertEqual(h.get("a"), b"c") - self.assertEqual(h.getlist("a"), [b"b", b"c"]) + assert h["a"] == b"c" + assert h.get("a") == b"c" + assert h.getlist("a") == [b"b", b"c"] def test_encode_utf8(self): h = Headers({"key": "\xa3"}, encoding="utf-8") key, val = dict(h).popitem() assert isinstance(key, bytes), key assert isinstance(val[0], bytes), val[0] - self.assertEqual(val[0], b"\xc2\xa3") + assert val[0] == b"\xc2\xa3" def test_encode_latin1(self): h = Headers({"key": "\xa3"}, encoding="latin1") key, val = dict(h).popitem() - self.assertEqual(val[0], b"\xa3") + assert val[0] == b"\xa3" def test_encode_multiple(self): h = Headers({"key": ["\xa3"]}, encoding="utf-8") key, val = dict(h).popitem() - self.assertEqual(val[0], b"\xc2\xa3") + assert val[0] == b"\xc2\xa3" def test_delete_and_contains(self): h = Headers() @@ -81,17 +80,17 @@ class HeadersTest(unittest.TestCase): h = Headers() olist = h.setdefault("X-Forwarded-For", "ip1") - self.assertEqual(h.getlist("X-Forwarded-For"), [b"ip1"]) + assert h.getlist("X-Forwarded-For") == [b"ip1"] assert h.getlist("X-Forwarded-For") is olist def test_iterables(self): idict = {"Content-Type": "text/html", "X-Forwarded-For": ["ip1", "ip2"]} h = Headers(idict) - self.assertDictEqual( - dict(h), - {b"Content-Type": [b"text/html"], b"X-Forwarded-For": [b"ip1", b"ip2"]}, - ) + assert dict(h) == { + b"Content-Type": [b"text/html"], + b"X-Forwarded-For": [b"ip1", b"ip2"], + } self.assertSortedEqual(h.keys(), [b"X-Forwarded-For", b"Content-Type"]) self.assertSortedEqual( h.items(), @@ -102,57 +101,57 @@ class HeadersTest(unittest.TestCase): def test_update(self): h = Headers() h.update({"Content-Type": "text/html", "X-Forwarded-For": ["ip1", "ip2"]}) - self.assertEqual(h.getlist("Content-Type"), [b"text/html"]) - self.assertEqual(h.getlist("X-Forwarded-For"), [b"ip1", b"ip2"]) + assert h.getlist("Content-Type") == [b"text/html"] + assert h.getlist("X-Forwarded-For") == [b"ip1", b"ip2"] def test_copy(self): h1 = Headers({"header1": ["value1", "value2"]}) h2 = copy.copy(h1) - self.assertEqual(h1, h2) - self.assertEqual(h1.getlist("header1"), h2.getlist("header1")) + assert h1 == h2 + assert h1.getlist("header1") == h2.getlist("header1") assert h1.getlist("header1") is not h2.getlist("header1") assert isinstance(h2, Headers) def test_appendlist(self): h1 = Headers({"header1": "value1"}) h1.appendlist("header1", "value3") - self.assertEqual(h1.getlist("header1"), [b"value1", b"value3"]) + assert h1.getlist("header1") == [b"value1", b"value3"] h1 = Headers() h1.appendlist("header1", "value1") h1.appendlist("header1", "value3") - self.assertEqual(h1.getlist("header1"), [b"value1", b"value3"]) + assert h1.getlist("header1") == [b"value1", b"value3"] def test_setlist(self): h1 = Headers({"header1": "value1"}) - self.assertEqual(h1.getlist("header1"), [b"value1"]) + assert h1.getlist("header1") == [b"value1"] h1.setlist("header1", [b"value2", b"value3"]) - self.assertEqual(h1.getlist("header1"), [b"value2", b"value3"]) + assert h1.getlist("header1") == [b"value2", b"value3"] def test_setlistdefault(self): h1 = Headers({"header1": "value1"}) h1.setlistdefault("header1", ["value2", "value3"]) h1.setlistdefault("header2", ["value2", "value3"]) - self.assertEqual(h1.getlist("header1"), [b"value1"]) - self.assertEqual(h1.getlist("header2"), [b"value2", b"value3"]) + assert h1.getlist("header1") == [b"value1"] + assert h1.getlist("header2") == [b"value2", b"value3"] def test_none_value(self): h1 = Headers() h1["foo"] = "bar" h1["foo"] = None h1.setdefault("foo", "bar") - self.assertEqual(h1.get("foo"), None) - self.assertEqual(h1.getlist("foo"), []) + assert h1.get("foo") is None + assert h1.getlist("foo") == [] def test_int_value(self): h1 = Headers({"hey": 5}) h1["foo"] = 1 h1.setdefault("bar", 2) h1.setlist("buz", [1, "dos", 3]) - self.assertEqual(h1.getlist("foo"), [b"1"]) - self.assertEqual(h1.getlist("bar"), [b"2"]) - self.assertEqual(h1.getlist("buz"), [b"1", b"dos", b"3"]) - self.assertEqual(h1.getlist("hey"), [b"5"]) + assert h1.getlist("foo") == [b"1"] + assert h1.getlist("bar") == [b"2"] + assert h1.getlist("buz") == [b"1", b"dos", b"3"] + assert h1.getlist("hey") == [b"5"] def test_invalid_value(self): with pytest.raises(TypeError, match="Unsupported value type"): From d442227fa74e414f4c7ac6baea6c3c4a1d938219 Mon Sep 17 00:00:00 2001 From: Andrey Rakhmatullin <wrar@wrar.name> Date: Sun, 9 Mar 2025 23:24:12 +0400 Subject: [PATCH 4/5] Converting tests to plain asserts, part 8. (#6711) --- tests/test_http_request.py | 608 +++++++++++++++----------------- tests/test_http_response.py | 316 ++++++++--------- tests/test_loader.py | 280 +++++++-------- tests/test_settings/__init__.py | 309 ++++++++-------- 4 files changed, 718 insertions(+), 795 deletions(-) diff --git a/tests/test_http_request.py b/tests/test_http_request.py index e5291157d..6bf0b8e3f 100644 --- a/tests/test_http_request.py +++ b/tests/test_http_request.py @@ -1,6 +1,5 @@ import json import re -import unittest import warnings import xmlrpc.client from typing import Any @@ -22,7 +21,7 @@ from scrapy.utils.httpobj import urlparse_cached from scrapy.utils.python import to_bytes, to_unicode -class RequestTest(unittest.TestCase): +class TestRequest: request_class = Request default_method = "GET" default_headers: dict[bytes, list[bytes]] = {} @@ -40,12 +39,12 @@ class RequestTest(unittest.TestCase): r = self.request_class("http://www.example.com") assert isinstance(r.url, str) - self.assertEqual(r.url, "http://www.example.com") - self.assertEqual(r.method, self.default_method) + assert r.url == "http://www.example.com" + assert r.method == self.default_method assert isinstance(r.headers, Headers) - self.assertEqual(r.headers, self.default_headers) - self.assertEqual(r.meta, self.default_meta) + assert r.headers == self.default_headers + assert r.meta == self.default_meta meta = {"lala": "lolo"} headers = {b"caca": b"coco"} @@ -54,9 +53,9 @@ class RequestTest(unittest.TestCase): ) assert r.meta is not meta - self.assertEqual(r.meta, meta) + assert r.meta == meta assert r.headers is not headers - self.assertEqual(r.headers[b"caca"], b"coco") + assert r.headers[b"caca"] == b"coco" def test_url_scheme(self): # This test passes by not raising any (ValueError) exception @@ -83,61 +82,61 @@ class RequestTest(unittest.TestCase): r = self.request_class(url=url, headers=headers) p = self.request_class(url=url, headers=r.headers) - self.assertEqual(r.headers, p.headers) - self.assertFalse(r.headers is headers) - self.assertFalse(p.headers is r.headers) + assert r.headers == p.headers + assert r.headers is not headers + assert p.headers is not r.headers # headers must not be unicode h = Headers({"key1": "val1", "key2": "val2"}) h["newkey"] = "newval" for k, v in h.items(): - self.assertIsInstance(k, bytes) + assert isinstance(k, bytes) for s in v: - self.assertIsInstance(s, bytes) + assert isinstance(s, bytes) def test_eq(self): url = "http://www.scrapy.org" r1 = self.request_class(url=url) r2 = self.request_class(url=url) - self.assertNotEqual(r1, r2) + assert r1 != r2 set_ = set() set_.add(r1) set_.add(r2) - self.assertEqual(len(set_), 2) + assert len(set_) == 2 def test_url(self): r = self.request_class(url="http://www.scrapy.org/path") - self.assertEqual(r.url, "http://www.scrapy.org/path") + assert r.url == "http://www.scrapy.org/path" def test_url_quoting(self): r = self.request_class(url="http://www.scrapy.org/blank%20space") - self.assertEqual(r.url, "http://www.scrapy.org/blank%20space") + assert r.url == "http://www.scrapy.org/blank%20space" r = self.request_class(url="http://www.scrapy.org/blank space") - self.assertEqual(r.url, "http://www.scrapy.org/blank%20space") + assert r.url == "http://www.scrapy.org/blank%20space" def test_url_encoding(self): r = self.request_class(url="http://www.scrapy.org/price/£") - self.assertEqual(r.url, "http://www.scrapy.org/price/%C2%A3") + assert r.url == "http://www.scrapy.org/price/%C2%A3" def test_url_encoding_other(self): # encoding affects only query part of URI, not path # path part should always be UTF-8 encoded before percent-escaping r = self.request_class(url="http://www.scrapy.org/price/£", encoding="utf-8") - self.assertEqual(r.url, "http://www.scrapy.org/price/%C2%A3") + assert r.url == "http://www.scrapy.org/price/%C2%A3" r = self.request_class(url="http://www.scrapy.org/price/£", encoding="latin1") - self.assertEqual(r.url, "http://www.scrapy.org/price/%C2%A3") + assert r.url == "http://www.scrapy.org/price/%C2%A3" def test_url_encoding_query(self): r1 = self.request_class(url="http://www.scrapy.org/price/£?unit=µ") - self.assertEqual(r1.url, "http://www.scrapy.org/price/%C2%A3?unit=%C2%B5") + assert r1.url == "http://www.scrapy.org/price/%C2%A3?unit=%C2%B5" # should be same as above r2 = self.request_class( url="http://www.scrapy.org/price/£?unit=µ", encoding="utf-8" ) - self.assertEqual(r2.url, "http://www.scrapy.org/price/%C2%A3?unit=%C2%B5") + assert r2.url == "http://www.scrapy.org/price/%C2%A3?unit=%C2%B5" def test_url_encoding_query_latin1(self): # encoding is used for encoding query-string before percent-escaping; @@ -145,7 +144,7 @@ class RequestTest(unittest.TestCase): r3 = self.request_class( url="http://www.scrapy.org/price/µ?currency=£", encoding="latin1" ) - self.assertEqual(r3.url, "http://www.scrapy.org/price/%C2%B5?currency=%A3") + assert r3.url == "http://www.scrapy.org/price/%C2%B5?currency=%A3" def test_url_encoding_nonutf8_untouched(self): # percent-escaping sequences that do not match valid UTF-8 sequences @@ -164,16 +163,16 @@ class RequestTest(unittest.TestCase): # "http://www.example.org/r%C3%A9sum%C3%A9.html", which is a different # URI from "http://www.example.org/r%E9sum%E9.html". r1 = self.request_class(url="http://www.scrapy.org/price/%a3") - self.assertEqual(r1.url, "http://www.scrapy.org/price/%a3") + assert r1.url == "http://www.scrapy.org/price/%a3" r2 = self.request_class(url="http://www.scrapy.org/r%C3%A9sum%C3%A9/%a3") - self.assertEqual(r2.url, "http://www.scrapy.org/r%C3%A9sum%C3%A9/%a3") + assert r2.url == "http://www.scrapy.org/r%C3%A9sum%C3%A9/%a3" r3 = self.request_class(url="http://www.scrapy.org/résumé/%a3") - self.assertEqual(r3.url, "http://www.scrapy.org/r%C3%A9sum%C3%A9/%a3") + assert r3.url == "http://www.scrapy.org/r%C3%A9sum%C3%A9/%a3" r4 = self.request_class(url="http://www.example.org/r%E9sum%E9.html") - self.assertEqual(r4.url, "http://www.example.org/r%E9sum%E9.html") + assert r4.url == "http://www.example.org/r%E9sum%E9.html" def test_body(self): r1 = self.request_class(url="http://www.example.com/") @@ -181,19 +180,19 @@ class RequestTest(unittest.TestCase): r2 = self.request_class(url="http://www.example.com/", body=b"") assert isinstance(r2.body, bytes) - self.assertEqual(r2.encoding, "utf-8") # default encoding + assert r2.encoding == "utf-8" # default encoding r3 = self.request_class( url="http://www.example.com/", body="Price: \xa3100", encoding="utf-8" ) assert isinstance(r3.body, bytes) - self.assertEqual(r3.body, b"Price: \xc2\xa3100") + assert r3.body == b"Price: \xc2\xa3100" r4 = self.request_class( url="http://www.example.com/", body="Price: \xa3100", encoding="latin1" ) assert isinstance(r4.body, bytes) - self.assertEqual(r4.body, b"Price: \xa3100") + assert r4.body == b"Price: \xa3100" def test_copy(self): """Test Request copy""" @@ -219,25 +218,25 @@ class RequestTest(unittest.TestCase): # make sure flags list is shallow copied assert r1.flags is not r2.flags, "flags must be a shallow copy, not identical" - self.assertEqual(r1.flags, r2.flags) + assert r1.flags == r2.flags # make sure cb_kwargs dict is shallow copied assert r1.cb_kwargs is not r2.cb_kwargs, ( "cb_kwargs must be a shallow copy, not identical" ) - self.assertEqual(r1.cb_kwargs, r2.cb_kwargs) + assert r1.cb_kwargs == r2.cb_kwargs # make sure meta dict is shallow copied assert r1.meta is not r2.meta, "meta must be a shallow copy, not identical" - self.assertEqual(r1.meta, r2.meta) + assert r1.meta == r2.meta # make sure headers attribute is shallow copied assert r1.headers is not r2.headers, ( "headers must be a shallow copy, not identical" ) - self.assertEqual(r1.headers, r2.headers) - self.assertEqual(r1.encoding, r2.encoding) - self.assertEqual(r1.dont_filter, r2.dont_filter) + assert r1.headers == r2.headers + assert r1.encoding == r2.encoding + assert r1.dont_filter == r2.dont_filter # Request.body can be identical since it's an immutable object (str) @@ -258,10 +257,10 @@ class RequestTest(unittest.TestCase): hdrs = Headers(r1.headers) hdrs[b"key"] = b"value" r2 = r1.replace(method="POST", body="New body", headers=hdrs) - self.assertEqual(r1.url, r2.url) - self.assertEqual((r1.method, r2.method), ("GET", "POST")) - self.assertEqual((r1.body, r2.body), (b"", b"New body")) - self.assertEqual((r1.headers, r2.headers), (self.default_headers, hdrs)) + assert r1.url == r2.url + assert (r1.method, r2.method) == ("GET", "POST") + assert (r1.body, r2.body) == (b"", b"New body") + assert (r1.headers, r2.headers) == (self.default_headers, hdrs) # Empty attributes (which may fail if not compared properly) r3 = self.request_class( @@ -270,9 +269,9 @@ class RequestTest(unittest.TestCase): r4 = r3.replace( url="http://www.example.com/2", body=b"", meta={}, dont_filter=False ) - self.assertEqual(r4.url, "http://www.example.com/2") - self.assertEqual(r4.body, b"") - self.assertEqual(r4.meta, {}) + assert r4.url == "http://www.example.com/2" + assert r4.body == b"" + assert r4.meta == {} assert r4.dont_filter is False def test_method_always_str(self): @@ -291,32 +290,32 @@ class RequestTest(unittest.TestCase): pass r1 = self.request_class("http://example.com") - self.assertIsNone(r1.callback) - self.assertIsNone(r1.errback) + assert r1.callback is None + assert r1.errback is None r2 = self.request_class("http://example.com", callback=a_function) - self.assertIs(r2.callback, a_function) - self.assertIsNone(r2.errback) + assert r2.callback is a_function + assert r2.errback is None r3 = self.request_class("http://example.com", errback=a_function) - self.assertIsNone(r3.callback) - self.assertIs(r3.errback, a_function) + assert r3.callback is None + assert r3.errback is a_function r4 = self.request_class( url="http://example.com", callback=a_function, errback=a_function, ) - self.assertIs(r4.callback, a_function) - self.assertIs(r4.errback, a_function) + assert r4.callback is a_function + assert r4.errback is a_function r5 = self.request_class( url="http://example.com", callback=NO_CALLBACK, errback=NO_CALLBACK, ) - self.assertIs(r5.callback, NO_CALLBACK) - self.assertIs(r5.errback, NO_CALLBACK) + assert r5.callback is NO_CALLBACK + assert r5.errback is NO_CALLBACK def test_callback_and_errback_type(self): with pytest.raises(TypeError): @@ -354,53 +353,46 @@ class RequestTest(unittest.TestCase): "2%3A15&comments=' --compressed" ) r = self.request_class.from_curl(curl_command) - self.assertEqual(r.method, "POST") - self.assertEqual(r.url, "http://httpbin.org/post") - self.assertEqual( - r.body, - b"custname=John+Smith&custtel=500&custemail=jsmith%40" + assert r.method == "POST" + assert r.url == "http://httpbin.org/post" + assert ( + r.body == b"custname=John+Smith&custtel=500&custemail=jsmith%40" b"example.org&size=small&topping=cheese&topping=onion" - b"&delivery=12%3A15&comments=", - ) - self.assertEqual( - r.cookies, - { - "_gauges_unique_year": "1", - "_gauges_unique": "1", - "_gauges_unique_month": "1", - "_gauges_unique_hour": "1", - "_gauges_unique_day": "1", - }, - ) - self.assertEqual( - r.headers, - { - b"Origin": [b"http://httpbin.org"], - b"Accept-Encoding": [b"gzip, deflate"], - b"Accept-Language": [b"en-US,en;q=0.9,ru;q=0.8,es;q=0.7"], - b"Upgrade-Insecure-Requests": [b"1"], - b"User-Agent": [ - b"Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537." - b"36 (KHTML, like Gecko) Ubuntu Chromium/62.0.3202" - b".75 Chrome/62.0.3202.75 Safari/537.36" - ], - b"Content-Type": [b"application /x-www-form-urlencoded"], - b"Accept": [ - b"text/html,application/xhtml+xml,application/xml;q=0." - b"9,image/webp,image/apng,*/*;q=0.8" - ], - b"Cache-Control": [b"max-age=0"], - b"Referer": [b"http://httpbin.org/forms/post"], - b"Connection": [b"keep-alive"], - }, + b"&delivery=12%3A15&comments=" ) + assert r.cookies == { + "_gauges_unique_year": "1", + "_gauges_unique": "1", + "_gauges_unique_month": "1", + "_gauges_unique_hour": "1", + "_gauges_unique_day": "1", + } + assert r.headers == { + b"Origin": [b"http://httpbin.org"], + b"Accept-Encoding": [b"gzip, deflate"], + b"Accept-Language": [b"en-US,en;q=0.9,ru;q=0.8,es;q=0.7"], + b"Upgrade-Insecure-Requests": [b"1"], + b"User-Agent": [ + b"Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537." + b"36 (KHTML, like Gecko) Ubuntu Chromium/62.0.3202" + b".75 Chrome/62.0.3202.75 Safari/537.36" + ], + b"Content-Type": [b"application /x-www-form-urlencoded"], + b"Accept": [ + b"text/html,application/xhtml+xml,application/xml;q=0." + b"9,image/webp,image/apng,*/*;q=0.8" + ], + b"Cache-Control": [b"max-age=0"], + b"Referer": [b"http://httpbin.org/forms/post"], + b"Connection": [b"keep-alive"], + } def test_from_curl_with_kwargs(self): r = self.request_class.from_curl( 'curl -X PATCH "http://example.org"', method="POST", meta={"key": "value"} ) - self.assertEqual(r.method, "POST") - self.assertEqual(r.meta, {"key": "value"}) + assert r.method == "POST" + assert r.meta == {"key": "value"} def test_from_curl_ignore_unknown_options(self): # By default: it works and ignores the unknown options: --foo and -z @@ -409,7 +401,7 @@ class RequestTest(unittest.TestCase): r = self.request_class.from_curl( 'curl -X DELETE "http://example.org" --foo -z', ) - self.assertEqual(r.method, "DELETE") + assert r.method == "DELETE" # If `ignore_unknown_options` is set to `False` it raises an error with # the unknown options: --foo and -z @@ -420,17 +412,17 @@ class RequestTest(unittest.TestCase): ) -class FormRequestTest(RequestTest): +class TestFormRequest(TestRequest): request_class = FormRequest def assertQueryEqual(self, first, second, msg=None): first = to_unicode(first).split("&") second = to_unicode(second).split("&") - return self.assertEqual(sorted(first), sorted(second), msg) + assert sorted(first) == sorted(second), msg def test_empty_formdata(self): r1 = self.request_class("http://www.example.com", formdata={}) - self.assertEqual(r1.body, b"") + assert r1.body == b"" def test_formdata_overrides_querystring(self): data = (("a", "one"), ("a", "two"), ("b", "2")) @@ -438,69 +430,61 @@ class FormRequestTest(RequestTest): "http://www.example.com/?a=0&b=1&c=3#fragment", method="GET", formdata=data ).url.split("#", maxsplit=1)[0] fs = _qs(self.request_class(url, method="GET", formdata=data)) - self.assertEqual(set(fs[b"a"]), {b"one", b"two"}) - self.assertEqual(fs[b"b"], [b"2"]) - self.assertIsNone(fs.get(b"c")) + assert set(fs[b"a"]) == {b"one", b"two"} + assert fs[b"b"] == [b"2"] + assert fs.get(b"c") is None data = {"a": "1", "b": "2"} fs = _qs( self.request_class("http://www.example.com/", method="GET", formdata=data) ) - self.assertEqual(fs[b"a"], [b"1"]) - self.assertEqual(fs[b"b"], [b"2"]) + assert fs[b"a"] == [b"1"] + assert fs[b"b"] == [b"2"] def test_default_encoding_bytes(self): # using default encoding (utf-8) data = {b"one": b"two", b"price": b"\xc2\xa3 100"} r2 = self.request_class("http://www.example.com", formdata=data) - self.assertEqual(r2.method, "POST") - self.assertEqual(r2.encoding, "utf-8") + assert r2.method == "POST" + assert r2.encoding == "utf-8" self.assertQueryEqual(r2.body, b"price=%C2%A3+100&one=two") - self.assertEqual( - r2.headers[b"Content-Type"], b"application/x-www-form-urlencoded" - ) + assert r2.headers[b"Content-Type"] == b"application/x-www-form-urlencoded" def test_default_encoding_textual_data(self): # using default encoding (utf-8) data = {"µ one": "two", "price": "£ 100"} r2 = self.request_class("http://www.example.com", formdata=data) - self.assertEqual(r2.method, "POST") - self.assertEqual(r2.encoding, "utf-8") + assert r2.method == "POST" + assert r2.encoding == "utf-8" self.assertQueryEqual(r2.body, b"price=%C2%A3+100&%C2%B5+one=two") - self.assertEqual( - r2.headers[b"Content-Type"], b"application/x-www-form-urlencoded" - ) + assert r2.headers[b"Content-Type"] == b"application/x-www-form-urlencoded" def test_default_encoding_mixed_data(self): # using default encoding (utf-8) data = {"\u00b5one": b"two", b"price\xc2\xa3": "\u00a3 100"} r2 = self.request_class("http://www.example.com", formdata=data) - self.assertEqual(r2.method, "POST") - self.assertEqual(r2.encoding, "utf-8") + assert r2.method == "POST" + assert r2.encoding == "utf-8" self.assertQueryEqual(r2.body, b"%C2%B5one=two&price%C2%A3=%C2%A3+100") - self.assertEqual( - r2.headers[b"Content-Type"], b"application/x-www-form-urlencoded" - ) + assert r2.headers[b"Content-Type"] == b"application/x-www-form-urlencoded" def test_custom_encoding_bytes(self): data = {b"\xb5 one": b"two", b"price": b"\xa3 100"} r2 = self.request_class( "http://www.example.com", formdata=data, encoding="latin1" ) - self.assertEqual(r2.method, "POST") - self.assertEqual(r2.encoding, "latin1") + assert r2.method == "POST" + assert r2.encoding == "latin1" self.assertQueryEqual(r2.body, b"price=%A3+100&%B5+one=two") - self.assertEqual( - r2.headers[b"Content-Type"], b"application/x-www-form-urlencoded" - ) + assert r2.headers[b"Content-Type"] == b"application/x-www-form-urlencoded" def test_custom_encoding_textual_data(self): data = {"price": "£ 100"} r3 = self.request_class( "http://www.example.com", formdata=data, encoding="latin1" ) - self.assertEqual(r3.encoding, "latin1") - self.assertEqual(r3.body, b"price=%A3+100") + assert r3.encoding == "latin1" + assert r3.body == b"price=%A3+100" def test_multi_key_values(self): # using multiples values for a single key @@ -523,16 +507,14 @@ class FormRequestTest(RequestTest): response, formdata={"one": ["two", "three"], "six": "seven"} ) - self.assertEqual(req.method, "POST") - self.assertEqual( - req.headers[b"Content-type"], b"application/x-www-form-urlencoded" - ) - self.assertEqual(req.url, "http://www.example.com/this/post.php") + assert req.method == "POST" + assert req.headers[b"Content-type"] == b"application/x-www-form-urlencoded" + assert req.url == "http://www.example.com/this/post.php" fs = _qs(req) - self.assertEqual(set(fs[b"test"]), {b"val1", b"val2"}) - self.assertEqual(set(fs[b"one"]), {b"two", b"three"}) - self.assertEqual(fs[b"test2"], [b"xxx"]) - self.assertEqual(fs[b"six"], [b"seven"]) + assert set(fs[b"test"]) == {b"val1", b"val2"} + assert set(fs[b"one"]) == {b"two", b"three"} + assert fs[b"test2"] == [b"xxx"] + assert fs[b"six"] == [b"seven"] def test_from_response_post_nonascii_bytes_utf8(self): response = _buildresponse( @@ -547,16 +529,14 @@ class FormRequestTest(RequestTest): response, formdata={"one": ["two", "three"], "six": "seven"} ) - self.assertEqual(req.method, "POST") - self.assertEqual( - req.headers[b"Content-type"], b"application/x-www-form-urlencoded" - ) - self.assertEqual(req.url, "http://www.example.com/this/post.php") + assert req.method == "POST" + assert req.headers[b"Content-type"] == b"application/x-www-form-urlencoded" + assert req.url == "http://www.example.com/this/post.php" fs = _qs(req, to_unicode=True) - self.assertEqual(set(fs["test £"]), {"val1", "val2"}) - self.assertEqual(set(fs["one"]), {"two", "three"}) - self.assertEqual(fs["test2"], ["xxx µ"]) - self.assertEqual(fs["six"], ["seven"]) + assert set(fs["test £"]) == {"val1", "val2"} + assert set(fs["one"]) == {"two", "three"} + assert fs["test2"] == ["xxx µ"] + assert fs["six"] == ["seven"] def test_from_response_post_nonascii_bytes_latin1(self): response = _buildresponse( @@ -572,16 +552,14 @@ class FormRequestTest(RequestTest): response, formdata={"one": ["two", "three"], "six": "seven"} ) - self.assertEqual(req.method, "POST") - self.assertEqual( - req.headers[b"Content-type"], b"application/x-www-form-urlencoded" - ) - self.assertEqual(req.url, "http://www.example.com/this/post.php") + assert req.method == "POST" + assert req.headers[b"Content-type"] == b"application/x-www-form-urlencoded" + assert req.url == "http://www.example.com/this/post.php" fs = _qs(req, to_unicode=True, encoding="latin1") - self.assertEqual(set(fs["test £"]), {"val1", "val2"}) - self.assertEqual(set(fs["one"]), {"two", "three"}) - self.assertEqual(fs["test2"], ["xxx µ"]) - self.assertEqual(fs["six"], ["seven"]) + assert set(fs["test £"]) == {"val1", "val2"} + assert set(fs["one"]) == {"two", "three"} + assert fs["test2"] == ["xxx µ"] + assert fs["six"] == ["seven"] def test_from_response_post_nonascii_unicode(self): response = _buildresponse( @@ -596,16 +574,14 @@ class FormRequestTest(RequestTest): response, formdata={"one": ["two", "three"], "six": "seven"} ) - self.assertEqual(req.method, "POST") - self.assertEqual( - req.headers[b"Content-type"], b"application/x-www-form-urlencoded" - ) - self.assertEqual(req.url, "http://www.example.com/this/post.php") + assert req.method == "POST" + assert req.headers[b"Content-type"] == b"application/x-www-form-urlencoded" + assert req.url == "http://www.example.com/this/post.php" fs = _qs(req, to_unicode=True) - self.assertEqual(set(fs["test £"]), {"val1", "val2"}) - self.assertEqual(set(fs["one"]), {"two", "three"}) - self.assertEqual(fs["test2"], ["xxx µ"]) - self.assertEqual(fs["six"], ["seven"]) + assert set(fs["test £"]) == {"val1", "val2"} + assert set(fs["one"]) == {"two", "three"} + assert fs["test2"] == ["xxx µ"] + assert fs["six"] == ["seven"] def test_from_response_duplicate_form_key(self): response = _buildresponse("<form></form>", url="http://www.example.com") @@ -614,8 +590,8 @@ class FormRequestTest(RequestTest): method="GET", formdata=(("foo", "bar"), ("foo", "baz")), ) - self.assertEqual(urlparse_cached(req).hostname, "www.example.com") - self.assertEqual(urlparse_cached(req).query, "foo=bar&foo=baz") + assert urlparse_cached(req).hostname == "www.example.com" + assert urlparse_cached(req).query == "foo=bar&foo=baz" def test_from_response_override_duplicate_form_key(self): response = _buildresponse( @@ -628,8 +604,8 @@ class FormRequestTest(RequestTest): response, formdata=(("two", "2"), ("two", "4")) ) fs = _qs(req) - self.assertEqual(fs[b"one"], [b"1"]) - self.assertEqual(fs[b"two"], [b"2", b"4"]) + assert fs[b"one"] == [b"1"] + assert fs[b"two"] == [b"2", b"4"] def test_from_response_extra_headers(self): response = _buildresponse( @@ -644,11 +620,9 @@ class FormRequestTest(RequestTest): formdata={"one": ["two", "three"], "six": "seven"}, headers={"Accept-Encoding": "gzip,deflate"}, ) - self.assertEqual(req.method, "POST") - self.assertEqual( - req.headers["Content-type"], b"application/x-www-form-urlencoded" - ) - self.assertEqual(req.headers["Accept-Encoding"], b"gzip,deflate") + assert req.method == "POST" + assert req.headers["Content-type"] == b"application/x-www-form-urlencoded" + assert req.headers["Accept-Encoding"] == b"gzip,deflate" def test_from_response_get(self): response = _buildresponse( @@ -662,14 +636,14 @@ class FormRequestTest(RequestTest): r1 = self.request_class.from_response( response, formdata={"one": ["two", "three"], "six": "seven"} ) - self.assertEqual(r1.method, "GET") - self.assertEqual(urlparse_cached(r1).hostname, "www.example.com") - self.assertEqual(urlparse_cached(r1).path, "/this/get.php") + assert r1.method == "GET" + assert urlparse_cached(r1).hostname == "www.example.com" + assert urlparse_cached(r1).path == "/this/get.php" fs = _qs(r1) - self.assertEqual(set(fs[b"test"]), {b"val1", b"val2"}) - self.assertEqual(set(fs[b"one"]), {b"two", b"three"}) - self.assertEqual(fs[b"test2"], [b"xxx"]) - self.assertEqual(fs[b"six"], [b"seven"]) + assert set(fs[b"test"]) == {b"val1", b"val2"} + assert set(fs[b"one"]) == {b"two", b"three"} + assert fs[b"test2"] == [b"xxx"] + assert fs[b"six"] == [b"seven"] def test_from_response_override_params(self): response = _buildresponse( @@ -680,8 +654,8 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(response, formdata={"two": "2"}) fs = _qs(req) - self.assertEqual(fs[b"one"], [b"1"]) - self.assertEqual(fs[b"two"], [b"2"]) + assert fs[b"one"] == [b"1"] + assert fs[b"two"] == [b"2"] def test_from_response_drop_params(self): response = _buildresponse( @@ -692,8 +666,8 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(response, formdata={"two": None}) fs = _qs(req) - self.assertEqual(fs[b"one"], [b"1"]) - self.assertNotIn(b"two", fs) + assert fs[b"one"] == [b"1"] + assert b"two" not in fs def test_from_response_override_method(self): response = _buildresponse( @@ -702,9 +676,9 @@ class FormRequestTest(RequestTest): </body></html>""" ) request = FormRequest.from_response(response) - self.assertEqual(request.method, "GET") + assert request.method == "GET" request = FormRequest.from_response(response, method="POST") - self.assertEqual(request.method, "POST") + assert request.method == "POST" def test_from_response_override_url(self): response = _buildresponse( @@ -713,11 +687,11 @@ class FormRequestTest(RequestTest): </body></html>""" ) request = FormRequest.from_response(response) - self.assertEqual(request.url, "http://example.com/app") + assert request.url == "http://example.com/app" request = FormRequest.from_response(response, url="http://foo.bar/absolute") - self.assertEqual(request.url, "http://foo.bar/absolute") + assert request.url == "http://foo.bar/absolute" request = FormRequest.from_response(response, url="/relative") - self.assertEqual(request.url, "http://example.com/relative") + assert request.url == "http://example.com/relative" def test_from_response_case_insensitive(self): response = _buildresponse( @@ -729,9 +703,9 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(response) fs = _qs(req) - self.assertEqual(fs[b"clickable1"], [b"clicked1"]) - self.assertFalse(b"i1" in fs, fs) # xpath in _get_inputs() - self.assertFalse(b"clickable2" in fs, fs) # xpath in _get_clickable() + assert fs[b"clickable1"] == [b"clicked1"] + assert b"i1" not in fs, fs # xpath in _get_inputs() + assert b"clickable2" not in fs, fs # xpath in _get_clickable() def test_from_response_submit_first_clickable(self): response = _buildresponse( @@ -744,10 +718,10 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(response, formdata={"two": "2"}) fs = _qs(req) - self.assertEqual(fs[b"clickable1"], [b"clicked1"]) - self.assertFalse(b"clickable2" in fs, fs) - self.assertEqual(fs[b"one"], [b"1"]) - self.assertEqual(fs[b"two"], [b"2"]) + assert fs[b"clickable1"] == [b"clicked1"] + assert b"clickable2" not in fs, fs + assert fs[b"one"] == [b"1"] + assert fs[b"two"] == [b"2"] def test_from_response_submit_not_first_clickable(self): response = _buildresponse( @@ -762,10 +736,10 @@ class FormRequestTest(RequestTest): response, formdata={"two": "2"}, clickdata={"name": "clickable2"} ) fs = _qs(req) - self.assertEqual(fs[b"clickable2"], [b"clicked2"]) - self.assertFalse(b"clickable1" in fs, fs) - self.assertEqual(fs[b"one"], [b"1"]) - self.assertEqual(fs[b"two"], [b"2"]) + assert fs[b"clickable2"] == [b"clicked2"] + assert b"clickable1" not in fs, fs + assert fs[b"one"] == [b"1"] + assert fs[b"two"] == [b"2"] def test_from_response_dont_submit_image_as_input(self): response = _buildresponse( @@ -777,7 +751,7 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(response, dont_click=True) fs = _qs(req) - self.assertEqual(fs, {b"i1": [b"i1v"]}) + assert fs == {b"i1": [b"i1v"]} def test_from_response_dont_submit_reset_as_input(self): response = _buildresponse( @@ -790,7 +764,7 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(response, dont_click=True) fs = _qs(req) - self.assertEqual(fs, {b"i1": [b"i1v"], b"i2": [b"i2v"]}) + assert fs == {b"i1": [b"i1v"], b"i2": [b"i2v"]} def test_from_response_clickdata_does_not_ignore_image(self): response = _buildresponse( @@ -801,7 +775,7 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(response) fs = _qs(req) - self.assertEqual(fs, {b"i1": [b"i1v"], b"i2": [b"i2v"]}) + assert fs == {b"i1": [b"i1v"], b"i2": [b"i2v"]} def test_from_response_multiple_clickdata(self): response = _buildresponse( @@ -816,9 +790,9 @@ class FormRequestTest(RequestTest): response, clickdata={"name": "clickable", "value": "clicked2"} ) fs = _qs(req) - self.assertEqual(fs[b"clickable"], [b"clicked2"]) - self.assertEqual(fs[b"one"], [b"clicked1"]) - self.assertEqual(fs[b"two"], [b"clicked2"]) + assert fs[b"clickable"] == [b"clicked2"] + assert fs[b"one"] == [b"clicked1"] + assert fs[b"two"] == [b"clicked2"] def test_from_response_unicode_clickdata(self): response = _buildresponse( @@ -833,7 +807,7 @@ class FormRequestTest(RequestTest): response, clickdata={"name": "price in \u00a3"} ) fs = _qs(req, to_unicode=True) - self.assertTrue(fs["price in \u00a3"]) + assert fs["price in \u00a3"] def test_from_response_unicode_clickdata_latin1(self): response = _buildresponse( @@ -849,7 +823,7 @@ class FormRequestTest(RequestTest): response, clickdata={"name": "price in \u00a5"} ) fs = _qs(req, to_unicode=True, encoding="latin1") - self.assertTrue(fs["price in \u00a5"]) + assert fs["price in \u00a5"] def test_from_response_multiple_forms_clickdata(self): response = _buildresponse( @@ -867,9 +841,9 @@ class FormRequestTest(RequestTest): response, formname="form2", clickdata={"name": "clickable"} ) fs = _qs(req) - self.assertEqual(fs[b"clickable"], [b"clicked2"]) - self.assertEqual(fs[b"field2"], [b"value2"]) - self.assertFalse(b"field1" in fs, fs) + assert fs[b"clickable"] == [b"clicked2"] + assert fs[b"field2"] == [b"value2"] + assert b"field1" not in fs, fs def test_from_response_override_clickable(self): response = _buildresponse( @@ -879,7 +853,7 @@ class FormRequestTest(RequestTest): response, formdata={"clickme": "two"}, clickdata={"name": "clickme"} ) fs = _qs(req) - self.assertEqual(fs[b"clickme"], [b"two"]) + assert fs[b"clickme"] == [b"two"] def test_from_response_dont_click(self): response = _buildresponse( @@ -892,8 +866,8 @@ class FormRequestTest(RequestTest): ) r1 = self.request_class.from_response(response, dont_click=True) fs = _qs(r1) - self.assertFalse(b"clickable1" in fs, fs) - self.assertFalse(b"clickable2" in fs, fs) + assert b"clickable1" not in fs, fs + assert b"clickable2" not in fs, fs def test_from_response_ambiguous_clickdata(self): response = _buildresponse( @@ -934,8 +908,8 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(response, clickdata={"nr": 1}) fs = _qs(req) - self.assertIn(b"clickable2", fs) - self.assertNotIn(b"clickable1", fs) + assert b"clickable2" in fs + assert b"clickable1" not in fs def test_from_response_invalid_nr_index_clickdata(self): response = _buildresponse( @@ -962,7 +936,7 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(response, formdata={"bar": "buz"}) fs = _qs(req) - self.assertEqual(fs, {b"foo": [b"xxx"], b"bar": [b"buz"]}) + assert fs == {b"foo": [b"xxx"], b"bar": [b"buz"]} def test_from_response_errors_formnumber(self): response = _buildresponse( @@ -983,12 +957,10 @@ class FormRequestTest(RequestTest): </form>""" ) r1 = self.request_class.from_response(response, formdata={"two": "3"}) - self.assertEqual(r1.method, "POST") - self.assertEqual( - r1.headers["Content-type"], b"application/x-www-form-urlencoded" - ) + assert r1.method == "POST" + assert r1.headers["Content-type"] == b"application/x-www-form-urlencoded" fs = _qs(r1) - self.assertEqual(fs, {b"one": [b"1"], b"two": [b"3"]}) + assert fs == {b"one": [b"1"], b"two": [b"3"]} def test_from_response_formname_exists(self): response = _buildresponse( @@ -1002,9 +974,9 @@ class FormRequestTest(RequestTest): </form>""" ) r1 = self.request_class.from_response(response, formname="form2") - self.assertEqual(r1.method, "POST") + assert r1.method == "POST" fs = _qs(r1) - self.assertEqual(fs, {b"four": [b"4"], b"three": [b"3"]}) + assert fs == {b"four": [b"4"], b"three": [b"3"]} def test_from_response_formname_nonexistent(self): response = _buildresponse( @@ -1016,9 +988,9 @@ class FormRequestTest(RequestTest): </form>""" ) r1 = self.request_class.from_response(response, formname="form3") - self.assertEqual(r1.method, "POST") + assert r1.method == "POST" fs = _qs(r1) - self.assertEqual(fs, {b"one": [b"1"]}) + assert fs == {b"one": [b"1"]} def test_from_response_formname_errors_formnumber(self): response = _buildresponse( @@ -1044,9 +1016,9 @@ class FormRequestTest(RequestTest): </form>""" ) r1 = self.request_class.from_response(response, formid="form2") - self.assertEqual(r1.method, "POST") + assert r1.method == "POST" fs = _qs(r1) - self.assertEqual(fs, {b"four": [b"4"], b"three": [b"3"]}) + assert fs == {b"four": [b"4"], b"three": [b"3"]} def test_from_response_formname_nonexistent_fallback_formid(self): response = _buildresponse( @@ -1062,9 +1034,9 @@ class FormRequestTest(RequestTest): r1 = self.request_class.from_response( response, formname="form3", formid="form2" ) - self.assertEqual(r1.method, "POST") + assert r1.method == "POST" fs = _qs(r1) - self.assertEqual(fs, {b"four": [b"4"], b"three": [b"3"]}) + assert fs == {b"four": [b"4"], b"three": [b"3"]} def test_from_response_formid_nonexistent(self): response = _buildresponse( @@ -1076,9 +1048,9 @@ class FormRequestTest(RequestTest): </form>""" ) r1 = self.request_class.from_response(response, formid="form3") - self.assertEqual(r1.method, "POST") + assert r1.method == "POST" fs = _qs(r1) - self.assertEqual(fs, {b"one": [b"1"]}) + assert fs == {b"one": [b"1"]} def test_from_response_formid_errors_formnumber(self): response = _buildresponse( @@ -1122,7 +1094,7 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(res) fs = _qs(req, to_unicode=True) - self.assertEqual(fs, {"i1": ["i1v2"], "i2": ["i2v1"], "i4": ["i4v2", "i4v3"]}) + assert fs == {"i1": ["i1v2"], "i2": ["i2v1"], "i4": ["i4v2", "i4v3"]} def test_from_response_radio(self): res = _buildresponse( @@ -1139,7 +1111,7 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(res) fs = _qs(req) - self.assertEqual(fs, {b"i1": [b"iv2"], b"i2": [b"on"]}) + assert fs == {b"i1": [b"iv2"], b"i2": [b"on"]} def test_from_response_checkbox(self): res = _buildresponse( @@ -1156,7 +1128,7 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(res) fs = _qs(req) - self.assertEqual(fs, {b"i1": [b"iv2"], b"i2": [b"on"]}) + assert fs == {b"i1": [b"iv2"], b"i2": [b"on"]} def test_from_response_input_text(self): res = _buildresponse( @@ -1170,7 +1142,7 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(res) fs = _qs(req) - self.assertEqual(fs, {b"i1": [b"i1v1"], b"i2": [b""], b"i4": [b"i4v1"]}) + assert fs == {b"i1": [b"i1v1"], b"i2": [b""], b"i4": [b"i4v1"]} def test_from_response_input_hidden(self): res = _buildresponse( @@ -1183,7 +1155,7 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(res) fs = _qs(req) - self.assertEqual(fs, {b"i1": [b"i1v1"], b"i2": [b""]}) + assert fs == {b"i1": [b"i1v1"], b"i2": [b""]} def test_from_response_input_textarea(self): res = _buildresponse( @@ -1196,7 +1168,7 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(res) fs = _qs(req) - self.assertEqual(fs, {b"i1": [b"i1v"], b"i2": [b""], b"i3": [b""]}) + assert fs == {b"i1": [b"i1v"], b"i2": [b""], b"i3": [b""]} def test_from_response_descendants(self): res = _buildresponse( @@ -1218,7 +1190,7 @@ class FormRequestTest(RequestTest): ) req = self.request_class.from_response(res) fs = _qs(req) - self.assertEqual(set(fs), {b"h2", b"i2", b"i1", b"i3", b"h1", b"i5", b"i4"}) + assert set(fs) == {b"h2", b"i2", b"i1", b"i3", b"h1", b"i5", b"i4"} def test_from_response_xpath(self): response = _buildresponse( @@ -1235,13 +1207,13 @@ class FormRequestTest(RequestTest): response, formxpath="//form[@action='post.php']" ) fs = _qs(r1) - self.assertEqual(fs[b"one"], [b"1"]) + assert fs[b"one"] == [b"1"] r1 = self.request_class.from_response( response, formxpath="//form/input[@name='four']" ) fs = _qs(r1) - self.assertEqual(fs[b"three"], [b"3"]) + assert fs[b"three"] == [b"3"] with pytest.raises(ValueError, match="No <form> element found with"): self.request_class.from_response( @@ -1254,7 +1226,7 @@ class FormRequestTest(RequestTest): response, formxpath="//form[@name='\u044a']" ) fs = _qs(r) - self.assertEqual(fs, {}) + assert not fs xpath = "//form[@name='\u03b1']" with pytest.raises(ValueError, match=re.escape(xpath)): @@ -1270,15 +1242,13 @@ class FormRequestTest(RequestTest): url="http://www.example.com/this/list.html", ) req = self.request_class.from_response(response) - self.assertEqual(req.method, "POST") - self.assertEqual( - req.headers["Content-type"], b"application/x-www-form-urlencoded" - ) - self.assertEqual(req.url, "http://www.example.com/this/post.php") + assert req.method == "POST" + assert req.headers["Content-type"] == b"application/x-www-form-urlencoded" + assert req.url == "http://www.example.com/this/post.php" fs = _qs(req) - self.assertEqual(fs[b"test1"], [b"val1"]) - self.assertEqual(fs[b"test2"], [b"val2"]) - self.assertEqual(fs[b"button1"], [b"submit1"]) + assert fs[b"test1"] == [b"val1"] + assert fs[b"test2"] == [b"val2"] + assert fs[b"button1"] == [b"submit1"] def test_from_response_button_notype(self): response = _buildresponse( @@ -1290,15 +1260,13 @@ class FormRequestTest(RequestTest): url="http://www.example.com/this/list.html", ) req = self.request_class.from_response(response) - self.assertEqual(req.method, "POST") - self.assertEqual( - req.headers["Content-type"], b"application/x-www-form-urlencoded" - ) - self.assertEqual(req.url, "http://www.example.com/this/post.php") + assert req.method == "POST" + assert req.headers["Content-type"] == b"application/x-www-form-urlencoded" + assert req.url == "http://www.example.com/this/post.php" fs = _qs(req) - self.assertEqual(fs[b"test1"], [b"val1"]) - self.assertEqual(fs[b"test2"], [b"val2"]) - self.assertEqual(fs[b"button1"], [b"submit1"]) + assert fs[b"test1"] == [b"val1"] + assert fs[b"test2"] == [b"val2"] + assert fs[b"button1"] == [b"submit1"] def test_from_response_submit_novalue(self): response = _buildresponse( @@ -1310,15 +1278,13 @@ class FormRequestTest(RequestTest): url="http://www.example.com/this/list.html", ) req = self.request_class.from_response(response) - self.assertEqual(req.method, "POST") - self.assertEqual( - req.headers["Content-type"], b"application/x-www-form-urlencoded" - ) - self.assertEqual(req.url, "http://www.example.com/this/post.php") + assert req.method == "POST" + assert req.headers["Content-type"] == b"application/x-www-form-urlencoded" + assert req.url == "http://www.example.com/this/post.php" fs = _qs(req) - self.assertEqual(fs[b"test1"], [b"val1"]) - self.assertEqual(fs[b"test2"], [b"val2"]) - self.assertEqual(fs[b"button1"], [b""]) + assert fs[b"test1"] == [b"val1"] + assert fs[b"test2"] == [b"val2"] + assert fs[b"button1"] == [b""] def test_from_response_button_novalue(self): response = _buildresponse( @@ -1330,15 +1296,13 @@ class FormRequestTest(RequestTest): url="http://www.example.com/this/list.html", ) req = self.request_class.from_response(response) - self.assertEqual(req.method, "POST") - self.assertEqual( - req.headers["Content-type"], b"application/x-www-form-urlencoded" - ) - self.assertEqual(req.url, "http://www.example.com/this/post.php") + assert req.method == "POST" + assert req.headers["Content-type"] == b"application/x-www-form-urlencoded" + assert req.url == "http://www.example.com/this/post.php" fs = _qs(req) - self.assertEqual(fs[b"test1"], [b"val1"]) - self.assertEqual(fs[b"test2"], [b"val2"]) - self.assertEqual(fs[b"button1"], [b""]) + assert fs[b"test1"] == [b"val1"] + assert fs[b"test2"] == [b"val2"] + assert fs[b"button1"] == [b""] def test_html_base_form_action(self): response = _buildresponse( @@ -1356,12 +1320,12 @@ class FormRequestTest(RequestTest): url="http://a.com/", ) req = self.request_class.from_response(response) - self.assertEqual(req.url, "http://b.com/test_form") + assert req.url == "http://b.com/test_form" def test_spaces_in_action(self): resp = _buildresponse('<body><form action=" path\n"></form></body>') req = self.request_class.from_response(resp) - self.assertEqual(req.url, "http://example.com/path") + assert req.url == "http://example.com/path" def test_from_response_css(self): response = _buildresponse( @@ -1378,11 +1342,11 @@ class FormRequestTest(RequestTest): response, formcss="form[action='post.php']" ) fs = _qs(r1) - self.assertEqual(fs[b"one"], [b"1"]) + assert fs[b"one"] == [b"1"] r1 = self.request_class.from_response(response, formcss="input[name='four']") fs = _qs(r1) - self.assertEqual(fs[b"three"], [b"3"]) + assert fs[b"three"] == [b"3"] with pytest.raises(ValueError, match="No <form> element found with"): self.request_class.from_response(response, formcss="input[name='abc']") @@ -1400,7 +1364,7 @@ class FormRequestTest(RequestTest): "</form>" ) r = self.request_class.from_response(response) - self.assertEqual(r.method, expected) + assert r.method == expected def test_form_response_with_invalid_formdata_type_error(self): """Test that a ValueError is raised for non-iterable and non-dict formdata input""" @@ -1464,23 +1428,20 @@ def _qs(req, encoding="utf-8", to_unicode=False): return parse_qs(uqs, True) -class XmlRpcRequestTest(RequestTest): +class TestXmlRpcRequest(TestRequest): request_class = XmlRpcRequest default_method = "POST" default_headers = {b"Content-Type": [b"text/xml"]} def _test_request(self, **kwargs): r = self.request_class("http://scrapytest.org/rpc2", **kwargs) - self.assertEqual(r.headers[b"Content-Type"], b"text/xml") - self.assertEqual( - r.body, - to_bytes( - xmlrpc.client.dumps(**kwargs), encoding=kwargs.get("encoding", "utf-8") - ), + assert r.headers[b"Content-Type"] == b"text/xml" + assert r.body == to_bytes( + xmlrpc.client.dumps(**kwargs), encoding=kwargs.get("encoding", "utf-8") ) - self.assertEqual(r.method, "POST") - self.assertEqual(r.encoding, kwargs.get("encoding", "utf-8")) - self.assertTrue(r.dont_filter, True) + assert r.method == "POST" + assert r.encoding == kwargs.get("encoding", "utf-8") + assert r.dont_filter, True def test_xmlrpc_dumps(self): self._test_request(params=("value",)) @@ -1497,7 +1458,7 @@ class XmlRpcRequestTest(RequestTest): self._test_request(params=("pas£",), encoding="latin1") -class JsonRequestTest(RequestTest): +class TestJsonRequest(TestRequest): request_class = JsonRequest default_method = "GET" default_headers = { @@ -1505,49 +1466,51 @@ class JsonRequestTest(RequestTest): b"Accept": [b"application/json, text/javascript, */*; q=0.01"], } - def setUp(self): + def setup_method(self): warnings.simplefilter("always") - super().setUp() + + def teardown_method(self): + warnings.resetwarnings() def test_data(self): r1 = self.request_class(url="http://www.example.com/") - self.assertEqual(r1.body, b"") + assert r1.body == b"" body = b"body" r2 = self.request_class(url="http://www.example.com/", body=body) - self.assertEqual(r2.body, body) + assert r2.body == body data = { "name": "value", } r3 = self.request_class(url="http://www.example.com/", data=data) - self.assertEqual(r3.body, to_bytes(json.dumps(data))) + assert r3.body == to_bytes(json.dumps(data)) # empty data r4 = self.request_class(url="http://www.example.com/", data=[]) - self.assertEqual(r4.body, to_bytes(json.dumps([]))) + assert r4.body == to_bytes(json.dumps([])) def test_data_method(self): # data is not passed r1 = self.request_class(url="http://www.example.com/") - self.assertEqual(r1.method, "GET") + assert r1.method == "GET" body = b"body" r2 = self.request_class(url="http://www.example.com/", body=body) - self.assertEqual(r2.method, "GET") + assert r2.method == "GET" data = { "name": "value", } r3 = self.request_class(url="http://www.example.com/", data=data) - self.assertEqual(r3.method, "POST") + assert r3.method == "POST" # method passed explicitly r4 = self.request_class(url="http://www.example.com/", data=data, method="GET") - self.assertEqual(r4.method, "GET") + assert r4.method == "GET" r5 = self.request_class(url="http://www.example.com/", data=[]) - self.assertEqual(r5.method, "POST") + assert r5.method == "POST" def test_body_data(self): """passing both body and data should result a warning""" @@ -1557,10 +1520,10 @@ class JsonRequestTest(RequestTest): } with warnings.catch_warnings(record=True) as _warnings: r5 = self.request_class(url="http://www.example.com/", body=body, data=data) - self.assertEqual(r5.body, body) - self.assertEqual(r5.method, "GET") - self.assertEqual(len(_warnings), 1) - self.assertIn("data will be ignored", str(_warnings[0].message)) + assert r5.body == body + assert r5.method == "GET" + assert len(_warnings) == 1 + assert "data will be ignored" in str(_warnings[0].message) def test_empty_body_data(self): """passing any body value and data should result a warning""" @@ -1569,10 +1532,10 @@ class JsonRequestTest(RequestTest): } with warnings.catch_warnings(record=True) as _warnings: r6 = self.request_class(url="http://www.example.com/", body=b"", data=data) - self.assertEqual(r6.body, b"") - self.assertEqual(r6.method, "GET") - self.assertEqual(len(_warnings), 1) - self.assertIn("data will be ignored", str(_warnings[0].message)) + assert r6.body == b"" + assert r6.method == "GET" + assert len(_warnings) == 1 + assert "data will be ignored" in str(_warnings[0].message) def test_body_none_data(self): data = { @@ -1580,15 +1543,15 @@ class JsonRequestTest(RequestTest): } with warnings.catch_warnings(record=True) as _warnings: r7 = self.request_class(url="http://www.example.com/", body=None, data=data) - self.assertEqual(r7.body, to_bytes(json.dumps(data))) - self.assertEqual(r7.method, "POST") - self.assertEqual(len(_warnings), 0) + assert r7.body == to_bytes(json.dumps(data)) + assert r7.method == "POST" + assert len(_warnings) == 0 def test_body_data_none(self): with warnings.catch_warnings(record=True) as _warnings: r8 = self.request_class(url="http://www.example.com/", body=None, data=None) - self.assertEqual(r8.method, "GET") - self.assertEqual(len(_warnings), 0) + assert r8.method == "GET" + assert len(_warnings) == 0 def test_dumps_sort_keys(self): """Test that sort_keys=True is passed to json.dumps by default""" @@ -1598,7 +1561,7 @@ class JsonRequestTest(RequestTest): with mock.patch("json.dumps", return_value=b"") as mock_dumps: self.request_class(url="http://www.example.com/", data=data) kwargs = mock_dumps.call_args[1] - self.assertEqual(kwargs["sort_keys"], True) + assert kwargs["sort_keys"] is True def test_dumps_kwargs(self): """Test that dumps_kwargs are passed to json.dumps""" @@ -1614,8 +1577,8 @@ class JsonRequestTest(RequestTest): url="http://www.example.com/", data=data, dumps_kwargs=dumps_kwargs ) kwargs = mock_dumps.call_args[1] - self.assertEqual(kwargs["ensure_ascii"], True) - self.assertEqual(kwargs["allow_nan"], True) + assert kwargs["ensure_ascii"] is True + assert kwargs["allow_nan"] is True def test_replace_data(self): data1 = { @@ -1626,7 +1589,7 @@ class JsonRequestTest(RequestTest): } r1 = self.request_class(url="http://www.example.com/", data=data1) r2 = r1.replace(data=data2) - self.assertEqual(r2.body, to_bytes(json.dumps(data2))) + assert r2.body == to_bytes(json.dumps(data2)) def test_replace_sort_keys(self): """Test that replace provides sort_keys=True to json.dumps""" @@ -1640,7 +1603,7 @@ class JsonRequestTest(RequestTest): with mock.patch("json.dumps", return_value=b"") as mock_dumps: r1.replace(data=data2) kwargs = mock_dumps.call_args[1] - self.assertEqual(kwargs["sort_keys"], True) + assert kwargs["sort_keys"] is True def test_replace_dumps_kwargs(self): """Test that dumps_kwargs are provided to json.dumps when replace is called""" @@ -1660,8 +1623,8 @@ class JsonRequestTest(RequestTest): with mock.patch("json.dumps", return_value=b"") as mock_dumps: r1.replace(data=data2) kwargs = mock_dumps.call_args[1] - self.assertEqual(kwargs["ensure_ascii"], True) - self.assertEqual(kwargs["allow_nan"], True) + assert kwargs["ensure_ascii"] is True + assert kwargs["allow_nan"] is True def test_replacement_both_body_and_data_warns(self): """Test that we get a warning if both body and data are passed""" @@ -1677,11 +1640,6 @@ class JsonRequestTest(RequestTest): with warnings.catch_warnings(record=True) as _warnings: r1.replace(data=data2, body=body2) - self.assertIn( - "Both body and data passed. data will be ignored", - str(_warnings[0].message), + assert "Both body and data passed. data will be ignored" in str( + _warnings[0].message ) - - def tearDown(self): - warnings.resetwarnings() - super().tearDown() diff --git a/tests/test_http_response.py b/tests/test_http_response.py index 5a943f084..fdef5adea 100644 --- a/tests/test_http_response.py +++ b/tests/test_http_response.py @@ -1,5 +1,4 @@ import codecs -import unittest from unittest import mock import pytest @@ -22,62 +21,56 @@ from scrapy.utils.python import to_unicode from tests import get_testdata -class BaseResponseTest(unittest.TestCase): +class TestResponseBase: response_class = Response def test_init(self): # Response requires url in the constructor with pytest.raises(TypeError): self.response_class() - self.assertTrue( - isinstance(self.response_class("http://example.com/"), self.response_class) + assert isinstance( + self.response_class("http://example.com/"), self.response_class ) with pytest.raises(TypeError): self.response_class(b"http://example.com") with pytest.raises(TypeError): self.response_class(url="http://example.com", body={}) # body can be str or None - self.assertTrue( - isinstance( - self.response_class("http://example.com/", body=b""), - self.response_class, - ) + assert isinstance( + self.response_class("http://example.com/", body=b""), + self.response_class, ) - self.assertTrue( - isinstance( - self.response_class("http://example.com/", body=b"body"), - self.response_class, - ) + assert isinstance( + self.response_class("http://example.com/", body=b"body"), + self.response_class, ) # test presence of all optional parameters - self.assertTrue( - isinstance( - self.response_class( - "http://example.com/", body=b"", headers={}, status=200 - ), - self.response_class, - ) + assert isinstance( + self.response_class( + "http://example.com/", body=b"", headers={}, status=200 + ), + self.response_class, ) r = self.response_class("http://www.example.com") assert isinstance(r.url, str) - self.assertEqual(r.url, "http://www.example.com") - self.assertEqual(r.status, 200) + assert r.url == "http://www.example.com" + assert r.status == 200 assert isinstance(r.headers, Headers) - self.assertEqual(r.headers, {}) + assert not r.headers headers = {"foo": "bar"} body = b"a body" r = self.response_class("http://www.example.com", headers=headers, body=body) assert r.headers is not headers - self.assertEqual(r.headers[b"foo"], b"bar") + assert r.headers[b"foo"] == b"bar" r = self.response_class("http://www.example.com", status=301) - self.assertEqual(r.status, 301) + assert r.status == 301 r = self.response_class("http://www.example.com", status="301") - self.assertEqual(r.status, 301) + assert r.status == 301 with pytest.raises(ValueError, match=r"invalid literal for int\(\)"): self.response_class("http://example.com", status="lala200") @@ -88,18 +81,18 @@ class BaseResponseTest(unittest.TestCase): r1.flags.append("cached") r2 = r1.copy() - self.assertEqual(r1.status, r2.status) - self.assertEqual(r1.body, r2.body) + assert r1.status == r2.status + assert r1.body == r2.body # make sure flags list is shallow copied assert r1.flags is not r2.flags, "flags must be a shallow copy, not identical" - self.assertEqual(r1.flags, r2.flags) + assert r1.flags == r2.flags # make sure headers attribute is shallow copied assert r1.headers is not r2.headers, ( "headers must be a shallow copy, not identical" ) - self.assertEqual(r1.headers, r2.headers) + assert r1.headers == r2.headers def test_copy_meta(self): req = Request("http://www.example.com") @@ -144,16 +137,16 @@ class BaseResponseTest(unittest.TestCase): r1 = self.response_class("http://www.example.com") r2 = r1.replace(status=301, body=b"New body", headers=hdrs) assert r1.body == b"" - self.assertEqual(r1.url, r2.url) - self.assertEqual((r1.status, r2.status), (200, 301)) - self.assertEqual((r1.body, r2.body), (b"", b"New body")) - self.assertEqual((r1.headers, r2.headers), ({}, hdrs)) + assert r1.url == r2.url + assert (r1.status, r2.status) == (200, 301) + assert (r1.body, r2.body) == (b"", b"New body") + assert (r1.headers, r2.headers) == ({}, hdrs) # Empty attributes (which may fail if not compared properly) r3 = self.response_class("http://www.example.com", flags=["cached"]) r4 = r3.replace(body=b"", flags=[]) - self.assertEqual(r4.body, b"") - self.assertEqual(r4.flags, []) + assert r4.body == b"" + assert not r4.flags def _assert_response_values(self, response, encoding, body): if isinstance(body, str): @@ -166,11 +159,11 @@ class BaseResponseTest(unittest.TestCase): assert isinstance(response.body, bytes) assert isinstance(response.text, str) self._assert_response_encoding(response, encoding) - self.assertEqual(response.body, body_bytes) - self.assertEqual(response.text, body_unicode) + assert response.body == body_bytes + assert response.text == body_unicode def _assert_response_encoding(self, response, encoding): - self.assertEqual(response.encoding, resolve_encoding(encoding)) + assert response.encoding == resolve_encoding(encoding) def test_immutable_attributes(self): r = self.response_class("http://example.com") @@ -183,7 +176,7 @@ class BaseResponseTest(unittest.TestCase): """Test urljoin shortcut (only for existence, since behavior equals urljoin)""" joined = self.response_class("http://www.example.com").urljoin("/test") absolute = "http://www.example.com/test" - self.assertEqual(joined, absolute) + assert joined == absolute def test_shortcut_attributes(self): r = self.response_class("http://example.com", body=b"hello") @@ -241,7 +234,7 @@ class BaseResponseTest(unittest.TestCase): def test_follow_flags(self): res = self.response_class("http://example.com/") fol = res.follow("http://example.com/", flags=["cached", "allowed"]) - self.assertEqual(fol.flags, ["cached", "allowed"]) + assert fol.flags == ["cached", "allowed"] # Response.follow_all @@ -276,7 +269,7 @@ class BaseResponseTest(unittest.TestCase): def test_follow_all_empty(self): r = self.response_class("http://example.com") - self.assertEqual([], list(r.follow_all([]))) + assert not list(r.follow_all([])) def test_follow_all_invalid(self): r = self.response_class("http://example.com") @@ -327,13 +320,13 @@ class BaseResponseTest(unittest.TestCase): ] fol = re.follow_all(urls, flags=["cached", "allowed"]) for req in fol: - self.assertEqual(req.flags, ["cached", "allowed"]) + assert req.flags == ["cached", "allowed"] def _assert_followed_url(self, follow_obj, target_url, response=None): if response is None: response = self._links_response() req = response.follow(follow_obj) - self.assertEqual(req.url, target_url) + assert req.url == target_url return req def _assert_followed_all_urls(self, follow_obj, target_urls, response=None): @@ -341,7 +334,7 @@ class BaseResponseTest(unittest.TestCase): response = self._links_response() followed = response.follow_all(follow_obj) for req, target in zip(followed, target_urls): - self.assertEqual(req.url, target) + assert req.url == target yield req def _links_response(self): @@ -353,7 +346,7 @@ class BaseResponseTest(unittest.TestCase): return self.response_class("http://example.com/index", body=body) -class TextResponseTest(BaseResponseTest): +class TestTextResponse(TestResponseBase): response_class = TextResponse def test_replace(self): @@ -365,10 +358,10 @@ class TextResponseTest(BaseResponseTest): r3 = r1.replace(url="http://www.example.com/other", encoding="latin1") assert isinstance(r2, self.response_class) - self.assertEqual(r2.url, "http://www.example.com/other") + assert r2.url == "http://www.example.com/other" self._assert_response_encoding(r2, "cp852") - self.assertEqual(r3.url, "http://www.example.com/other") - self.assertEqual(r3._declared_encoding(), "latin1") + assert r3.url == "http://www.example.com/other" + assert r3._declared_encoding() == "latin1" def test_unicode_url(self): # instantiate with unicode url without encoding (should set default encoding) @@ -382,21 +375,21 @@ class TextResponseTest(BaseResponseTest): resp = self.response_class( url="http://www.example.com/price/\xa3", encoding="utf-8" ) - self.assertEqual(resp.url, to_unicode(b"http://www.example.com/price/\xc2\xa3")) + assert resp.url == to_unicode(b"http://www.example.com/price/\xc2\xa3") resp = self.response_class( url="http://www.example.com/price/\xa3", encoding="latin-1" ) - self.assertEqual(resp.url, "http://www.example.com/price/\xa3") + assert resp.url == "http://www.example.com/price/\xa3" resp = self.response_class( "http://www.example.com/price/\xa3", headers={"Content-type": ["text/html; charset=utf-8"]}, ) - self.assertEqual(resp.url, to_unicode(b"http://www.example.com/price/\xc2\xa3")) + assert resp.url == to_unicode(b"http://www.example.com/price/\xc2\xa3") resp = self.response_class( "http://www.example.com/price/\xa3", headers={"Content-type": ["text/html; charset=iso-8859-1"]}, ) - self.assertEqual(resp.url, "http://www.example.com/price/\xa3") + assert resp.url == "http://www.example.com/price/\xa3" def test_unicode_body(self): unicode_string = ( @@ -412,8 +405,8 @@ class TextResponseTest(BaseResponseTest): ) # check response.text - self.assertTrue(isinstance(r1.text, str)) - self.assertEqual(r1.text, unicode_string) + assert isinstance(r1.text, str) + assert r1.text == unicode_string def test_encoding(self): r1 = self.response_class( @@ -458,18 +451,18 @@ class TextResponseTest(BaseResponseTest): }, ) - self.assertEqual(r1._headers_encoding(), "utf-8") - self.assertEqual(r2._headers_encoding(), None) - self.assertEqual(r2._declared_encoding(), "utf-8") + assert r1._headers_encoding() == "utf-8" + assert r2._headers_encoding() is None + assert r2._declared_encoding() == "utf-8" self._assert_response_encoding(r2, "utf-8") - self.assertEqual(r3._headers_encoding(), "cp1252") - self.assertEqual(r3._declared_encoding(), "cp1252") - self.assertEqual(r4._headers_encoding(), None) - self.assertEqual(r5._headers_encoding(), None) - self.assertEqual(r8._headers_encoding(), "cp1251") - self.assertEqual(r9._headers_encoding(), None) - self.assertEqual(r8._declared_encoding(), "utf-8") - self.assertEqual(r9._declared_encoding(), None) + assert r3._headers_encoding() == "cp1252" + assert r3._declared_encoding() == "cp1252" + assert r4._headers_encoding() is None + assert r5._headers_encoding() is None + assert r8._headers_encoding() == "cp1251" + assert r9._headers_encoding() is None + assert r8._declared_encoding() == "utf-8" + assert r9._declared_encoding() is None self._assert_response_encoding(r5, "utf-8") self._assert_response_encoding(r8, "utf-8") self._assert_response_encoding(r9, "cp1252") @@ -493,7 +486,7 @@ class TextResponseTest(BaseResponseTest): headers={"Content-type": ["text/html; charset=UNKNOWN"]}, body=b"\xc2\xa3", ) - self.assertEqual(r._declared_encoding(), None) + assert r._declared_encoding() is None self._assert_response_values(r, "utf-8", "\xa3") def test_utf16(self): @@ -511,14 +504,11 @@ class TextResponseTest(BaseResponseTest): headers={"Content-type": ["text/html; charset=utf-8"]}, body=b"\xef\xbb\xbfWORD\xe3\xab", ) - self.assertEqual(r6.encoding, "utf-8") - self.assertIn( - r6.text, - { - "WORD\ufffd\ufffd", # w3lib < 1.19.0 - "WORD\ufffd", # w3lib >= 1.19.0 - }, - ) + assert r6.encoding == "utf-8" + assert r6.text in { + "WORD\ufffd\ufffd", # w3lib < 1.19.0 + "WORD\ufffd", # w3lib >= 1.19.0 + } def test_bom_is_removed_from_body(self): # Inferring encoding from body also cache decoded body as sideeffect, @@ -532,21 +522,21 @@ class TextResponseTest(BaseResponseTest): # Test response without content-type and BOM encoding response = self.response_class(url, body=body) - self.assertEqual(response.encoding, "utf-8") - self.assertEqual(response.text, "WORD") + assert response.encoding == "utf-8" + assert response.text == "WORD" response = self.response_class(url, body=body) - self.assertEqual(response.text, "WORD") - self.assertEqual(response.encoding, "utf-8") + assert response.text == "WORD" + assert response.encoding == "utf-8" # Body caching sideeffect isn't triggered when encoding is declared in # content-type header but BOM still need to be removed from decoded # body response = self.response_class(url, headers=headers, body=body) - self.assertEqual(response.encoding, "utf-8") - self.assertEqual(response.text, "WORD") + assert response.encoding == "utf-8" + assert response.text == "WORD" response = self.response_class(url, headers=headers, body=body) - self.assertEqual(response.text, "WORD") - self.assertEqual(response.encoding, "utf-8") + assert response.text == "WORD" + assert response.encoding == "utf-8" def test_replace_wrong_encoding(self): """Test invalid chars are replaced properly""" @@ -577,49 +567,47 @@ class TextResponseTest(BaseResponseTest): body = b"<html><head><title>Some page" response = self.response_class("http://www.example.com", body=body) - self.assertIsInstance(response.selector, Selector) - self.assertEqual(response.selector.type, "html") - self.assertIs(response.selector, response.selector) # property is cached - self.assertIs(response.selector.response, response) + assert isinstance(response.selector, Selector) + assert response.selector.type == "html" + assert response.selector is response.selector # property is cached + assert response.selector.response is response - self.assertEqual( - response.selector.xpath("//title/text()").getall(), ["Some page"] - ) - self.assertEqual(response.selector.css("title::text").getall(), ["Some page"]) - self.assertEqual(response.selector.re("Some (.*)"), ["page"]) + assert response.selector.xpath("//title/text()").getall() == ["Some page"] + assert response.selector.css("title::text").getall() == ["Some page"] + assert response.selector.re("Some (.*)") == ["page"] def test_selector_shortcuts(self): body = b"Some page" response = self.response_class("http://www.example.com", body=body) - self.assertEqual( - response.xpath("//title/text()").getall(), - response.selector.xpath("//title/text()").getall(), + assert ( + response.xpath("//title/text()").getall() + == response.selector.xpath("//title/text()").getall() ) - self.assertEqual( - response.css("title::text").getall(), - response.selector.css("title::text").getall(), + assert ( + response.css("title::text").getall() + == response.selector.css("title::text").getall() ) def test_selector_shortcuts_kwargs(self): body = b'Some page

A nice paragraph.

' response = self.response_class("http://www.example.com", body=body) - self.assertEqual( + assert ( response.xpath( "normalize-space(//p[@class=$pclass])", pclass="content" - ).getall(), - response.xpath('normalize-space(//p[@class="content"])').getall(), + ).getall() + == response.xpath('normalize-space(//p[@class="content"])').getall() ) - self.assertEqual( + assert ( response.xpath( "//title[count(following::p[@class=$pclass])=$pcount]/text()", pclass="content", pcount=1, - ).getall(), - response.xpath( + ).getall() + == response.xpath( '//title[count(following::p[@class="content"])=1]/text()' - ).getall(), + ).getall() ) def test_urljoin_with_base_url(self): @@ -629,21 +617,21 @@ class TextResponseTest(BaseResponseTest): "/test" ) absolute = "https://example.net/test" - self.assertEqual(joined, absolute) + assert joined == absolute body = b'' joined = self.response_class("http://www.example.com", body=body).urljoin( "test" ) absolute = "http://www.example.com/test" - self.assertEqual(joined, absolute) + assert joined == absolute body = b'' joined = self.response_class("http://www.example.com", body=body).urljoin( "test" ) absolute = "http://www.example.com/elsewhere/test" - self.assertEqual(joined, absolute) + assert joined == absolute def test_follow_selector(self): resp = self._links_response() @@ -728,7 +716,7 @@ class TextResponseTest(BaseResponseTest): "http://example.com/foo?%D0%BF%D1%80%D0%B8%D0%B2%D0%B5%D1%82", response=resp1, ) - self.assertEqual(req.encoding, "utf8") + assert req.encoding == "utf8" resp2 = self.response_class( "http://example.com", @@ -742,12 +730,12 @@ class TextResponseTest(BaseResponseTest): "http://example.com/foo?%EF%F0%E8%E2%E5%F2", response=resp2, ) - self.assertEqual(req.encoding, "cp1251") + assert req.encoding == "cp1251" def test_follow_flags(self): res = self.response_class("http://example.com/") fol = res.follow("http://example.com/", flags=["cached", "allowed"]) - self.assertEqual(fol.flags, ["cached", "allowed"]) + assert fol.flags == ["cached", "allowed"] def test_follow_all_flags(self): re = self.response_class("http://www.example.com/") @@ -758,7 +746,7 @@ class TextResponseTest(BaseResponseTest): ] fol = re.follow_all(urls, flags=["cached", "allowed"]) for req in fol: - self.assertEqual(req.flags, ["cached", "allowed"]) + assert req.flags == ["cached", "allowed"] def test_follow_all_css(self): expected = [ @@ -767,7 +755,7 @@ class TextResponseTest(BaseResponseTest): ] response = self._links_response() extracted = [r.url for r in response.follow_all(css='a[href*="example.com"]')] - self.assertEqual(expected, extracted) + assert expected == extracted def test_follow_all_css_skip_invalid(self): expected = [ @@ -777,9 +765,9 @@ class TextResponseTest(BaseResponseTest): ] response = self._links_response_no_href() extracted1 = [r.url for r in response.follow_all(css=".pagination a")] - self.assertEqual(expected, extracted1) + assert expected == extracted1 extracted2 = [r.url for r in response.follow_all(response.css(".pagination a"))] - self.assertEqual(expected, extracted2) + assert expected == extracted2 def test_follow_all_xpath(self): expected = [ @@ -788,7 +776,7 @@ class TextResponseTest(BaseResponseTest): ] response = self._links_response() extracted = response.follow_all(xpath='//a[contains(@href, "example.com")]') - self.assertEqual(expected, [r.url for r in extracted]) + assert expected == [r.url for r in extracted] def test_follow_all_xpath_skip_invalid(self): expected = [ @@ -800,12 +788,12 @@ class TextResponseTest(BaseResponseTest): extracted1 = [ r.url for r in response.follow_all(xpath='//div[@id="pagination"]/a') ] - self.assertEqual(expected, extracted1) + assert expected == extracted1 extracted2 = [ r.url for r in response.follow_all(response.xpath('//div[@id="pagination"]/a')) ] - self.assertEqual(expected, extracted2) + assert expected == extracted2 def test_follow_all_too_many_arguments(self): response = self._links_response() @@ -820,7 +808,7 @@ class TextResponseTest(BaseResponseTest): def test_json_response(self): json_body = b"""{"ip": "109.187.217.200"}""" json_response = self.response_class("http://www.example.com", body=json_body) - self.assertEqual(json_response.json(), {"ip": "109.187.217.200"}) + assert json_response.json() == {"ip": "109.187.217.200"} text_body = b"""text""" text_response = self.response_class("http://www.example.com", body=text_body) @@ -842,7 +830,7 @@ class TextResponseTest(BaseResponseTest): mock_json.assert_called_once_with(json_body) -class HtmlResponseTest(TextResponseTest): +class TestHtmlResponse(TestTextResponse): response_class = HtmlResponse def test_html_encoding(self): @@ -883,7 +871,7 @@ class HtmlResponseTest(TextResponseTest): self._assert_response_values(r1, "gb2312", body) -class XmlResponseTest(TextResponseTest): +class TestXmlResponse(TestTextResponse): response_class = XmlResponse def test_xml_encoding(self): @@ -917,20 +905,20 @@ class XmlResponseTest(TextResponseTest): body = b'value' response = self.response_class("http://www.example.com", body=body) - self.assertIsInstance(response.selector, Selector) - self.assertEqual(response.selector.type, "xml") - self.assertIs(response.selector, response.selector) # property is cached - self.assertIs(response.selector.response, response) + assert isinstance(response.selector, Selector) + assert response.selector.type == "xml" + assert response.selector is response.selector # property is cached + assert response.selector.response is response - self.assertEqual(response.selector.xpath("//elem/text()").getall(), ["value"]) + assert response.selector.xpath("//elem/text()").getall() == ["value"] def test_selector_shortcuts(self): body = b'value' response = self.response_class("http://www.example.com", body=body) - self.assertEqual( - response.xpath("//elem/text()").getall(), - response.selector.xpath("//elem/text()").getall(), + assert ( + response.xpath("//elem/text()").getall() + == response.selector.xpath("//elem/text()").getall() ) def test_selector_shortcuts_kwargs(self): @@ -940,21 +928,21 @@ class XmlResponseTest(TextResponseTest): """ response = self.response_class("http://www.example.com", body=body) - self.assertEqual( + assert ( response.xpath( "//s:elem/text()", namespaces={"s": "http://scrapy.org"} - ).getall(), - response.selector.xpath( + ).getall() + == response.selector.xpath( "//s:elem/text()", namespaces={"s": "http://scrapy.org"} - ).getall(), + ).getall() ) response.selector.register_namespace("s2", "http://scrapy.org") - self.assertEqual( + assert ( response.xpath( "//s1:elem/text()", namespaces={"s1": "http://scrapy.org"} - ).getall(), - response.selector.xpath("//s2:elem/text()").getall(), + ).getall() + == response.selector.xpath("//s2:elem/text()").getall() ) @@ -968,7 +956,7 @@ class CustomResponse(TextResponse): super().__init__(*args, **kwargs) -class CustomResponseTest(TextResponseTest): +class TestCustomResponse(TestTextResponse): response_class = CustomResponse def test_copy(self): @@ -981,11 +969,11 @@ class CustomResponseTest(TextResponseTest): lost="lost", ) r2 = r1.copy() - self.assertIsInstance(r2, self.response_class) - self.assertEqual(r1.foo, r2.foo) - self.assertEqual(r1.bar, r2.bar) - self.assertEqual(r1.lost, "lost") - self.assertIsNone(r2.lost) + assert isinstance(r2, self.response_class) + assert r1.foo == r2.foo + assert r1.bar == r2.bar + assert r1.lost == "lost" + assert r2.lost is None def test_replace(self): super().test_replace() @@ -998,31 +986,31 @@ class CustomResponseTest(TextResponseTest): ) r2 = r1.replace(foo="new-foo", bar="new-bar", lost="new-lost") - self.assertIsInstance(r2, self.response_class) - self.assertEqual(r1.foo, "foo") - self.assertEqual(r1.bar, "bar") - self.assertEqual(r1.lost, "lost") - self.assertEqual(r2.foo, "new-foo") - self.assertEqual(r2.bar, "new-bar") - self.assertEqual(r2.lost, "new-lost") + assert isinstance(r2, self.response_class) + assert r1.foo == "foo" + assert r1.bar == "bar" + assert r1.lost == "lost" + assert r2.foo == "new-foo" + assert r2.bar == "new-bar" + assert r2.lost == "new-lost" r3 = r1.replace(foo="new-foo", bar="new-bar") - self.assertIsInstance(r3, self.response_class) - self.assertEqual(r1.foo, "foo") - self.assertEqual(r1.bar, "bar") - self.assertEqual(r1.lost, "lost") - self.assertEqual(r3.foo, "new-foo") - self.assertEqual(r3.bar, "new-bar") - self.assertIsNone(r3.lost) + assert isinstance(r3, self.response_class) + assert r1.foo == "foo" + assert r1.bar == "bar" + assert r1.lost == "lost" + assert r3.foo == "new-foo" + assert r3.bar == "new-bar" + assert r3.lost is None r4 = r1.replace(foo="new-foo") - self.assertIsInstance(r4, self.response_class) - self.assertEqual(r1.foo, "foo") - self.assertEqual(r1.bar, "bar") - self.assertEqual(r1.lost, "lost") - self.assertEqual(r4.foo, "new-foo") - self.assertEqual(r4.bar, "bar") - self.assertIsNone(r4.lost) + assert isinstance(r4, self.response_class) + assert r1.foo == "foo" + assert r1.bar == "bar" + assert r1.lost == "lost" + assert r4.foo == "new-foo" + assert r4.bar == "bar" + assert r4.lost is None with pytest.raises( TypeError, diff --git a/tests/test_loader.py b/tests/test_loader.py index 1a933bb8d..224158e7f 100644 --- a/tests/test_loader.py +++ b/tests/test_loader.py @@ -1,7 +1,6 @@ from __future__ import annotations import dataclasses -import unittest import attr import pytest @@ -67,7 +66,7 @@ def processor_with_args(value, other=None, loader_context=None): return value -class BasicItemLoaderTest(unittest.TestCase): +class TestBasicItemLoader: def test_add_value_on_unknown_field(self): il = ProcessorItemLoader() with pytest.raises(KeyError): @@ -80,14 +79,14 @@ class BasicItemLoaderTest(unittest.TestCase): il.add_value("name", "marta") item = il.load_item() assert item is i - self.assertEqual(item["summary"], ["lala"]) - self.assertEqual(item["name"], ["marta"]) + assert item["summary"] == ["lala"] + assert item["name"] == ["marta"] def test_load_item_using_custom_loader(self): il = ProcessorItemLoader() il.add_value("name", "marta") item = il.load_item() - self.assertEqual(item["name"], ["Marta"]) + assert item["name"] == ["Marta"] class InitializationTestMixin: @@ -98,16 +97,16 @@ class InitializationTestMixin: input_item = self.item_class(name="foo") il = ItemLoader(item=input_item) loaded_item = il.load_item() - self.assertIsInstance(loaded_item, self.item_class) - self.assertEqual(ItemAdapter(loaded_item).asdict(), {"name": ["foo"]}) + assert isinstance(loaded_item, self.item_class) + assert ItemAdapter(loaded_item).asdict() == {"name": ["foo"]} def test_keep_list(self): """Loaded item should contain values from the initial item""" input_item = self.item_class(name=["foo", "bar"]) il = ItemLoader(item=input_item) loaded_item = il.load_item() - self.assertIsInstance(loaded_item, self.item_class) - self.assertEqual(ItemAdapter(loaded_item).asdict(), {"name": ["foo", "bar"]}) + assert isinstance(loaded_item, self.item_class) + assert ItemAdapter(loaded_item).asdict() == {"name": ["foo", "bar"]} def test_add_value_singlevalue_singlevalue(self): """Values added after initialization should be appended""" @@ -115,8 +114,8 @@ class InitializationTestMixin: il = ItemLoader(item=input_item) il.add_value("name", "bar") loaded_item = il.load_item() - self.assertIsInstance(loaded_item, self.item_class) - self.assertEqual(ItemAdapter(loaded_item).asdict(), {"name": ["foo", "bar"]}) + assert isinstance(loaded_item, self.item_class) + assert ItemAdapter(loaded_item).asdict() == {"name": ["foo", "bar"]} def test_add_value_singlevalue_list(self): """Values added after initialization should be appended""" @@ -124,10 +123,8 @@ class InitializationTestMixin: il = ItemLoader(item=input_item) il.add_value("name", ["item", "loader"]) loaded_item = il.load_item() - self.assertIsInstance(loaded_item, self.item_class) - self.assertEqual( - ItemAdapter(loaded_item).asdict(), {"name": ["foo", "item", "loader"]} - ) + assert isinstance(loaded_item, self.item_class) + assert ItemAdapter(loaded_item).asdict() == {"name": ["foo", "item", "loader"]} def test_add_value_list_singlevalue(self): """Values added after initialization should be appended""" @@ -135,10 +132,8 @@ class InitializationTestMixin: il = ItemLoader(item=input_item) il.add_value("name", "qwerty") loaded_item = il.load_item() - self.assertIsInstance(loaded_item, self.item_class) - self.assertEqual( - ItemAdapter(loaded_item).asdict(), {"name": ["foo", "bar", "qwerty"]} - ) + assert isinstance(loaded_item, self.item_class) + assert ItemAdapter(loaded_item).asdict() == {"name": ["foo", "bar", "qwerty"]} def test_add_value_list_list(self): """Values added after initialization should be appended""" @@ -146,56 +141,55 @@ class InitializationTestMixin: il = ItemLoader(item=input_item) il.add_value("name", ["item", "loader"]) loaded_item = il.load_item() - self.assertIsInstance(loaded_item, self.item_class) - self.assertEqual( - ItemAdapter(loaded_item).asdict(), - {"name": ["foo", "bar", "item", "loader"]}, - ) + assert isinstance(loaded_item, self.item_class) + assert ItemAdapter(loaded_item).asdict() == { + "name": ["foo", "bar", "item", "loader"] + } def test_get_output_value_singlevalue(self): """Getting output value must not remove value from item""" input_item = self.item_class(name="foo") il = ItemLoader(item=input_item) - self.assertEqual(il.get_output_value("name"), ["foo"]) + assert il.get_output_value("name") == ["foo"] loaded_item = il.load_item() - self.assertIsInstance(loaded_item, self.item_class) - self.assertEqual(ItemAdapter(loaded_item).asdict(), {"name": ["foo"]}) + assert isinstance(loaded_item, self.item_class) + assert ItemAdapter(loaded_item).asdict() == {"name": ["foo"]} def test_get_output_value_list(self): """Getting output value must not remove value from item""" input_item = self.item_class(name=["foo", "bar"]) il = ItemLoader(item=input_item) - self.assertEqual(il.get_output_value("name"), ["foo", "bar"]) + assert il.get_output_value("name") == ["foo", "bar"] loaded_item = il.load_item() - self.assertIsInstance(loaded_item, self.item_class) - self.assertEqual(ItemAdapter(loaded_item).asdict(), {"name": ["foo", "bar"]}) + assert isinstance(loaded_item, self.item_class) + assert ItemAdapter(loaded_item).asdict() == {"name": ["foo", "bar"]} def test_values_single(self): """Values from initial item must be added to loader._values""" input_item = self.item_class(name="foo") il = ItemLoader(item=input_item) - self.assertEqual(il._values.get("name"), ["foo"]) + assert il._values.get("name") == ["foo"] def test_values_list(self): """Values from initial item must be added to loader._values""" input_item = self.item_class(name=["foo", "bar"]) il = ItemLoader(item=input_item) - self.assertEqual(il._values.get("name"), ["foo", "bar"]) + assert il._values.get("name") == ["foo", "bar"] -class InitializationFromDictTest(InitializationTestMixin, unittest.TestCase): +class TestInitializationFromDict(InitializationTestMixin): item_class = dict -class InitializationFromItemTest(InitializationTestMixin, unittest.TestCase): +class TestInitializationFromItem(InitializationTestMixin): item_class = NameItem -class InitializationFromAttrsItemTest(InitializationTestMixin, unittest.TestCase): +class TestInitializationFromAttrsItem(InitializationTestMixin): item_class = AttrsNameItem -class InitializationFromDataClassTest(InitializationTestMixin, unittest.TestCase): +class TestInitializationFromDataClass(InitializationTestMixin): item_class = NameDataClass @@ -212,7 +206,7 @@ class NoInputReprocessingItemLoader(BaseNoInputReprocessingLoader): default_item_class = NoInputReprocessingItem -class NoInputReprocessingFromItemTest(unittest.TestCase): +class TestNoInputReprocessingFromItem: """ Loaders initialized from loaded items must not reprocess fields (Item instances) """ @@ -220,41 +214,41 @@ class NoInputReprocessingFromItemTest(unittest.TestCase): def test_avoid_reprocessing_with_initial_values_single(self): il = NoInputReprocessingItemLoader(item=NoInputReprocessingItem(title="foo")) il_loaded = il.load_item() - self.assertEqual(il_loaded, {"title": "foo"}) - self.assertEqual( - NoInputReprocessingItemLoader(item=il_loaded).load_item(), {"title": "foo"} - ) + assert il_loaded == {"title": "foo"} + assert NoInputReprocessingItemLoader(item=il_loaded).load_item() == { + "title": "foo" + } def test_avoid_reprocessing_with_initial_values_list(self): il = NoInputReprocessingItemLoader( item=NoInputReprocessingItem(title=["foo", "bar"]) ) il_loaded = il.load_item() - self.assertEqual(il_loaded, {"title": "foo"}) - self.assertEqual( - NoInputReprocessingItemLoader(item=il_loaded).load_item(), {"title": "foo"} - ) + assert il_loaded == {"title": "foo"} + assert NoInputReprocessingItemLoader(item=il_loaded).load_item() == { + "title": "foo" + } def test_avoid_reprocessing_without_initial_values_single(self): il = NoInputReprocessingItemLoader() il.add_value("title", "FOO") il_loaded = il.load_item() - self.assertEqual(il_loaded, {"title": "FOO"}) - self.assertEqual( - NoInputReprocessingItemLoader(item=il_loaded).load_item(), {"title": "FOO"} - ) + assert il_loaded == {"title": "FOO"} + assert NoInputReprocessingItemLoader(item=il_loaded).load_item() == { + "title": "FOO" + } def test_avoid_reprocessing_without_initial_values_list(self): il = NoInputReprocessingItemLoader() il.add_value("title", ["foo", "bar"]) il_loaded = il.load_item() - self.assertEqual(il_loaded, {"title": "FOO"}) - self.assertEqual( - NoInputReprocessingItemLoader(item=il_loaded).load_item(), {"title": "FOO"} - ) + assert il_loaded == {"title": "FOO"} + assert NoInputReprocessingItemLoader(item=il_loaded).load_item() == { + "title": "FOO" + } -class TestOutputProcessorItem(unittest.TestCase): +class TestOutputProcessorItem: def test_output_processor(self): class TempItem(Item): temp = Field() @@ -270,11 +264,11 @@ class TestOutputProcessorItem(unittest.TestCase): loader = TempLoader() item = loader.load_item() - self.assertIsInstance(item, TempItem) - self.assertEqual(dict(item), {"temp": 0.3}) + assert isinstance(item, TempItem) + assert dict(item) == {"temp": 0.3} -class SelectortemLoaderTest(unittest.TestCase): +class TestSelectortemLoader: response = HtmlResponse( url="", encoding="utf-8", @@ -292,7 +286,7 @@ class SelectortemLoaderTest(unittest.TestCase): def test_init_method(self): l = ProcessorItemLoader() - self.assertEqual(l.selector, None) + assert l.selector is None def test_init_method_errors(self): l = ProcessorItemLoader() @@ -312,150 +306,149 @@ class SelectortemLoaderTest(unittest.TestCase): def test_init_method_with_selector(self): sel = Selector(text="
marta
") l = ProcessorItemLoader(selector=sel) - self.assertIs(l.selector, sel) + assert l.selector is sel l.add_xpath("name", "//div/text()") - self.assertEqual(l.get_output_value("name"), ["Marta"]) + assert l.get_output_value("name") == ["Marta"] def test_init_method_with_selector_css(self): sel = Selector(text="
marta
") l = ProcessorItemLoader(selector=sel) - self.assertIs(l.selector, sel) + assert l.selector is sel l.add_css("name", "div::text") - self.assertEqual(l.get_output_value("name"), ["Marta"]) + assert l.get_output_value("name") == ["Marta"] def test_init_method_with_base_response(self): """Selector should be None after initialization""" response = Response("https://scrapy.org") l = ProcessorItemLoader(response=response) - self.assertIs(l.selector, None) + assert l.selector is None def test_init_method_with_response(self): l = ProcessorItemLoader(response=self.response) - self.assertTrue(l.selector) + assert l.selector l.add_xpath("name", "//div/text()") - self.assertEqual(l.get_output_value("name"), ["Marta"]) + assert l.get_output_value("name") == ["Marta"] def test_init_method_with_response_css(self): l = ProcessorItemLoader(response=self.response) - self.assertTrue(l.selector) + assert l.selector l.add_css("name", "div::text") - self.assertEqual(l.get_output_value("name"), ["Marta"]) + assert l.get_output_value("name") == ["Marta"] l.add_css("url", "a::attr(href)") - self.assertEqual(l.get_output_value("url"), ["http://www.scrapy.org"]) + assert l.get_output_value("url") == ["http://www.scrapy.org"] # combining/accumulating CSS selectors and XPath expressions l.add_xpath("name", "//div/text()") - self.assertEqual(l.get_output_value("name"), ["Marta", "Marta"]) + assert l.get_output_value("name") == ["Marta", "Marta"] l.add_xpath("url", "//img/@src") - self.assertEqual( - l.get_output_value("url"), ["http://www.scrapy.org", "/images/logo.png"] - ) + assert l.get_output_value("url") == [ + "http://www.scrapy.org", + "/images/logo.png", + ] def test_add_xpath_re(self): l = ProcessorItemLoader(response=self.response) l.add_xpath("name", "//div/text()", re="ma") - self.assertEqual(l.get_output_value("name"), ["Ma"]) + assert l.get_output_value("name") == ["Ma"] def test_replace_xpath(self): l = ProcessorItemLoader(response=self.response) - self.assertTrue(l.selector) + assert l.selector l.add_xpath("name", "//div/text()") - self.assertEqual(l.get_output_value("name"), ["Marta"]) + assert l.get_output_value("name") == ["Marta"] l.replace_xpath("name", "//p/text()") - self.assertEqual(l.get_output_value("name"), ["Paragraph"]) + assert l.get_output_value("name") == ["Paragraph"] l.replace_xpath("name", ["//p/text()", "//div/text()"]) - self.assertEqual(l.get_output_value("name"), ["Paragraph", "Marta"]) + assert l.get_output_value("name") == ["Paragraph", "Marta"] def test_get_xpath(self): l = ProcessorItemLoader(response=self.response) - self.assertEqual(l.get_xpath("//p/text()"), ["paragraph"]) - self.assertEqual(l.get_xpath("//p/text()", TakeFirst()), "paragraph") - self.assertEqual(l.get_xpath("//p/text()", TakeFirst(), re="pa"), "pa") + assert l.get_xpath("//p/text()") == ["paragraph"] + assert l.get_xpath("//p/text()", TakeFirst()) == "paragraph" + assert l.get_xpath("//p/text()", TakeFirst(), re="pa") == "pa" - self.assertEqual( - l.get_xpath(["//p/text()", "//div/text()"]), ["paragraph", "marta"] - ) + assert l.get_xpath(["//p/text()", "//div/text()"]) == ["paragraph", "marta"] def test_replace_xpath_multi_fields(self): l = ProcessorItemLoader(response=self.response) l.add_xpath(None, "//div/text()", TakeFirst(), lambda x: {"name": x}) - self.assertEqual(l.get_output_value("name"), ["Marta"]) + assert l.get_output_value("name") == ["Marta"] l.replace_xpath(None, "//p/text()", TakeFirst(), lambda x: {"name": x}) - self.assertEqual(l.get_output_value("name"), ["Paragraph"]) + assert l.get_output_value("name") == ["Paragraph"] def test_replace_xpath_re(self): l = ProcessorItemLoader(response=self.response) - self.assertTrue(l.selector) + assert l.selector l.add_xpath("name", "//div/text()") - self.assertEqual(l.get_output_value("name"), ["Marta"]) + assert l.get_output_value("name") == ["Marta"] l.replace_xpath("name", "//div/text()", re="ma") - self.assertEqual(l.get_output_value("name"), ["Ma"]) + assert l.get_output_value("name") == ["Ma"] def test_add_css_re(self): l = ProcessorItemLoader(response=self.response) l.add_css("name", "div::text", re="ma") - self.assertEqual(l.get_output_value("name"), ["Ma"]) + assert l.get_output_value("name") == ["Ma"] l.add_css("url", "a::attr(href)", re="http://(.+)") - self.assertEqual(l.get_output_value("url"), ["www.scrapy.org"]) + assert l.get_output_value("url") == ["www.scrapy.org"] def test_replace_css(self): l = ProcessorItemLoader(response=self.response) - self.assertTrue(l.selector) + assert l.selector l.add_css("name", "div::text") - self.assertEqual(l.get_output_value("name"), ["Marta"]) + assert l.get_output_value("name") == ["Marta"] l.replace_css("name", "p::text") - self.assertEqual(l.get_output_value("name"), ["Paragraph"]) + assert l.get_output_value("name") == ["Paragraph"] l.replace_css("name", ["p::text", "div::text"]) - self.assertEqual(l.get_output_value("name"), ["Paragraph", "Marta"]) + assert l.get_output_value("name") == ["Paragraph", "Marta"] l.add_css("url", "a::attr(href)", re="http://(.+)") - self.assertEqual(l.get_output_value("url"), ["www.scrapy.org"]) + assert l.get_output_value("url") == ["www.scrapy.org"] l.replace_css("url", "img::attr(src)") - self.assertEqual(l.get_output_value("url"), ["/images/logo.png"]) + assert l.get_output_value("url") == ["/images/logo.png"] def test_get_css(self): l = ProcessorItemLoader(response=self.response) - self.assertEqual(l.get_css("p::text"), ["paragraph"]) - self.assertEqual(l.get_css("p::text", TakeFirst()), "paragraph") - self.assertEqual(l.get_css("p::text", TakeFirst(), re="pa"), "pa") + assert l.get_css("p::text") == ["paragraph"] + assert l.get_css("p::text", TakeFirst()) == "paragraph" + assert l.get_css("p::text", TakeFirst(), re="pa") == "pa" - self.assertEqual(l.get_css(["p::text", "div::text"]), ["paragraph", "marta"]) - self.assertEqual( - l.get_css(["a::attr(href)", "img::attr(src)"]), - ["http://www.scrapy.org", "/images/logo.png"], - ) + assert l.get_css(["p::text", "div::text"]) == ["paragraph", "marta"] + assert l.get_css(["a::attr(href)", "img::attr(src)"]) == [ + "http://www.scrapy.org", + "/images/logo.png", + ] def test_replace_css_multi_fields(self): l = ProcessorItemLoader(response=self.response) l.add_css(None, "div::text", TakeFirst(), lambda x: {"name": x}) - self.assertEqual(l.get_output_value("name"), ["Marta"]) + assert l.get_output_value("name") == ["Marta"] l.replace_css(None, "p::text", TakeFirst(), lambda x: {"name": x}) - self.assertEqual(l.get_output_value("name"), ["Paragraph"]) + assert l.get_output_value("name") == ["Paragraph"] l.add_css(None, "a::attr(href)", TakeFirst(), lambda x: {"url": x}) - self.assertEqual(l.get_output_value("url"), ["http://www.scrapy.org"]) + assert l.get_output_value("url") == ["http://www.scrapy.org"] l.replace_css(None, "img::attr(src)", TakeFirst(), lambda x: {"url": x}) - self.assertEqual(l.get_output_value("url"), ["/images/logo.png"]) + assert l.get_output_value("url") == ["/images/logo.png"] def test_replace_css_re(self): l = ProcessorItemLoader(response=self.response) - self.assertTrue(l.selector) + assert l.selector l.add_css("url", "a::attr(href)") - self.assertEqual(l.get_output_value("url"), ["http://www.scrapy.org"]) + assert l.get_output_value("url") == ["http://www.scrapy.org"] l.replace_css("url", "a::attr(href)", re=r"http://www\.(.+)") - self.assertEqual(l.get_output_value("url"), ["scrapy.org"]) + assert l.get_output_value("url") == ["scrapy.org"] -class SubselectorLoaderTest(unittest.TestCase): +class TestSubselectorLoader: response = HtmlResponse( url="", encoding="utf-8", @@ -483,17 +476,13 @@ class SubselectorLoaderTest(unittest.TestCase): nl.add_css("name_div", "#id") nl.add_value("name_value", nl.selector.xpath('div[@id = "id"]/text()').getall()) - self.assertEqual(l.get_output_value("name"), ["marta"]) - self.assertEqual(l.get_output_value("name_div"), ['
marta
']) - self.assertEqual(l.get_output_value("name_value"), ["marta"]) + assert l.get_output_value("name") == ["marta"] + assert l.get_output_value("name_div") == ['
marta
'] + assert l.get_output_value("name_value") == ["marta"] - self.assertEqual(l.get_output_value("name"), nl.get_output_value("name")) - self.assertEqual( - l.get_output_value("name_div"), nl.get_output_value("name_div") - ) - self.assertEqual( - l.get_output_value("name_value"), nl.get_output_value("name_value") - ) + assert l.get_output_value("name") == nl.get_output_value("name") + assert l.get_output_value("name_div") == nl.get_output_value("name_div") + assert l.get_output_value("name_value") == nl.get_output_value("name_value") def test_nested_css(self): l = NestedItemLoader(response=self.response) @@ -502,17 +491,13 @@ class SubselectorLoaderTest(unittest.TestCase): nl.add_css("name_div", "#id") nl.add_value("name_value", nl.selector.xpath('div[@id = "id"]/text()').getall()) - self.assertEqual(l.get_output_value("name"), ["marta"]) - self.assertEqual(l.get_output_value("name_div"), ['
marta
']) - self.assertEqual(l.get_output_value("name_value"), ["marta"]) + assert l.get_output_value("name") == ["marta"] + assert l.get_output_value("name_div") == ['
marta
'] + assert l.get_output_value("name_value") == ["marta"] - self.assertEqual(l.get_output_value("name"), nl.get_output_value("name")) - self.assertEqual( - l.get_output_value("name_div"), nl.get_output_value("name_div") - ) - self.assertEqual( - l.get_output_value("name_value"), nl.get_output_value("name_value") - ) + assert l.get_output_value("name") == nl.get_output_value("name") + assert l.get_output_value("name_div") == nl.get_output_value("name_div") + assert l.get_output_value("name_value") == nl.get_output_value("name_value") def test_nested_replace(self): l = NestedItemLoader(response=self.response) @@ -520,11 +505,11 @@ class SubselectorLoaderTest(unittest.TestCase): nl2 = nl1.nested_xpath("a") l.add_xpath("url", "//footer/a/@href") - self.assertEqual(l.get_output_value("url"), ["http://www.scrapy.org"]) + assert l.get_output_value("url") == ["http://www.scrapy.org"] nl1.replace_xpath("url", "img/@src") - self.assertEqual(l.get_output_value("url"), ["/images/logo.png"]) + assert l.get_output_value("url") == ["/images/logo.png"] nl2.replace_xpath("url", "@href") - self.assertEqual(l.get_output_value("url"), ["http://www.scrapy.org"]) + assert l.get_output_value("url") == ["http://www.scrapy.org"] def test_nested_ordering(self): l = NestedItemLoader(response=self.response) @@ -536,15 +521,12 @@ class SubselectorLoaderTest(unittest.TestCase): nl2.add_xpath("url", "text()") l.add_xpath("url", "//footer/a/@href") - self.assertEqual( - l.get_output_value("url"), - [ - "/images/logo.png", - "http://www.scrapy.org", - "homepage", - "http://www.scrapy.org", - ], - ) + assert l.get_output_value("url") == [ + "/images/logo.png", + "http://www.scrapy.org", + "homepage", + "http://www.scrapy.org", + ] def test_nested_load_item(self): l = NestedItemLoader(response=self.response) @@ -561,9 +543,9 @@ class SubselectorLoaderTest(unittest.TestCase): assert item is nl1.item assert item is nl2.item - self.assertEqual(item["name"], ["marta"]) - self.assertEqual(item["url"], ["http://www.scrapy.org"]) - self.assertEqual(item["image"], ["/images/logo.png"]) + assert item["name"] == ["marta"] + assert item["url"] == ["http://www.scrapy.org"] + assert item["image"] == ["/images/logo.png"] # Functions as processors @@ -588,9 +570,9 @@ class FunctionProcessorItemLoader(ItemLoader): default_item_class = FunctionProcessorItem -class FunctionProcessorTestCase(unittest.TestCase): +class TestFunctionProcessor: def test_processor_defined_in_item(self): lo = FunctionProcessorItemLoader() lo.add_value("foo", " bar ") lo.add_value("foo", [" asdf ", " qwerty "]) - self.assertEqual(dict(lo.load_item()), {"foo": ["BAR", "ASDF", "QWERTY"]}) + assert dict(lo.load_item()) == {"foo": ["BAR", "ASDF", "QWERTY"]} diff --git a/tests/test_settings/__init__.py b/tests/test_settings/__init__.py index 909b365a9..d7d900546 100644 --- a/tests/test_settings/__init__.py +++ b/tests/test_settings/__init__.py @@ -1,4 +1,6 @@ -import unittest +# pylint: disable=unsubscriptable-object,unsupported-membership-test,use-implicit-booleaness-not-comparison +# (too many false positives) + from unittest import mock import pytest @@ -14,31 +16,31 @@ from scrapy.settings import ( from . import default_settings -class SettingsGlobalFuncsTest(unittest.TestCase): +class TestSettingsGlobalFuncs: def test_get_settings_priority(self): for prio_str, prio_num in SETTINGS_PRIORITIES.items(): - self.assertEqual(get_settings_priority(prio_str), prio_num) - self.assertEqual(get_settings_priority(99), 99) + assert get_settings_priority(prio_str) == prio_num + assert get_settings_priority(99) == 99 -class SettingsAttributeTest(unittest.TestCase): - def setUp(self): +class TestSettingsAttribute: + def setup_method(self): self.attribute = SettingsAttribute("value", 10) def test_set_greater_priority(self): self.attribute.set("value2", 20) - self.assertEqual(self.attribute.value, "value2") - self.assertEqual(self.attribute.priority, 20) + assert self.attribute.value == "value2" + assert self.attribute.priority == 20 def test_set_equal_priority(self): self.attribute.set("value2", 10) - self.assertEqual(self.attribute.value, "value2") - self.assertEqual(self.attribute.priority, 10) + assert self.attribute.value == "value2" + assert self.attribute.priority == 10 def test_set_less_priority(self): self.attribute.set("value2", 0) - self.assertEqual(self.attribute.value, "value") - self.assertEqual(self.attribute.priority, 10) + assert self.attribute.value == "value" + assert self.attribute.priority == 10 def test_overwrite_basesettings(self): original_dict = {"one": 10, "two": 20} @@ -47,61 +49,59 @@ class SettingsAttributeTest(unittest.TestCase): new_dict = {"three": 11, "four": 21} attribute.set(new_dict, 10) - self.assertIsInstance(attribute.value, BaseSettings) - self.assertCountEqual(attribute.value, new_dict) - self.assertCountEqual(original_settings, original_dict) + assert isinstance(attribute.value, BaseSettings) + assert set(attribute.value) == set(new_dict) + assert set(original_settings) == set(original_dict) new_settings = BaseSettings({"five": 12}, 0) attribute.set(new_settings, 0) # Insufficient priority - self.assertCountEqual(attribute.value, new_dict) + assert set(attribute.value) == set(new_dict) attribute.set(new_settings, 10) - self.assertCountEqual(attribute.value, new_settings) + assert set(attribute.value) == set(new_settings) def test_repr(self): - self.assertEqual( - repr(self.attribute), "" - ) + assert repr(self.attribute) == "" -class BaseSettingsTest(unittest.TestCase): - def setUp(self): +class TestBaseSettings: + def setup_method(self): self.settings = BaseSettings() def test_setdefault_not_existing_value(self): settings = BaseSettings() value = settings.setdefault("TEST_OPTION", "value") - self.assertEqual(settings["TEST_OPTION"], "value") - self.assertEqual(value, "value") - self.assertIsNotNone(value) + assert settings["TEST_OPTION"] == "value" + assert value == "value" + assert value is not None def test_setdefault_existing_value(self): settings = BaseSettings({"TEST_OPTION": "value"}) value = settings.setdefault("TEST_OPTION", None) - self.assertEqual(settings["TEST_OPTION"], "value") - self.assertEqual(value, "value") + assert settings["TEST_OPTION"] == "value" + assert value == "value" def test_set_new_attribute(self): self.settings.set("TEST_OPTION", "value", 0) - self.assertIn("TEST_OPTION", self.settings.attributes) + assert "TEST_OPTION" in self.settings.attributes attr = self.settings.attributes["TEST_OPTION"] - self.assertIsInstance(attr, SettingsAttribute) - self.assertEqual(attr.value, "value") - self.assertEqual(attr.priority, 0) + assert isinstance(attr, SettingsAttribute) + assert attr.value == "value" + assert attr.priority == 0 def test_set_settingsattribute(self): myattr = SettingsAttribute(0, 30) # Note priority 30 self.settings.set("TEST_ATTR", myattr, 10) - self.assertEqual(self.settings.get("TEST_ATTR"), 0) - self.assertEqual(self.settings.getpriority("TEST_ATTR"), 30) + assert self.settings.get("TEST_ATTR") == 0 + assert self.settings.getpriority("TEST_ATTR") == 30 def test_set_instance_identity_on_update(self): attr = SettingsAttribute("value", 0) self.settings.attributes = {"TEST_OPTION": attr} self.settings.set("TEST_OPTION", "othervalue", 10) - self.assertIn("TEST_OPTION", self.settings.attributes) - self.assertIs(attr, self.settings.attributes["TEST_OPTION"]) + assert "TEST_OPTION" in self.settings.attributes + assert attr is self.settings.attributes["TEST_OPTION"] def test_set_calls_settings_attributes_methods_on_update(self): attr = SettingsAttribute("value", 10) @@ -114,7 +114,7 @@ class BaseSettingsTest(unittest.TestCase): for priority in (0, 10, 20): self.settings.set("TEST_OPTION", "othervalue", priority) mock_set.assert_called_once_with("othervalue", priority) - self.assertFalse(mock_setattr.called) + assert not mock_setattr.called mock_set.reset_mock() mock_setattr.reset_mock() @@ -122,19 +122,19 @@ class BaseSettingsTest(unittest.TestCase): settings = BaseSettings() settings.set("key", "a", "default") settings["key"] = "b" - self.assertEqual(settings["key"], "b") - self.assertEqual(settings.getpriority("key"), 20) + assert settings["key"] == "b" + assert settings.getpriority("key") == 20 settings["key"] = "c" - self.assertEqual(settings["key"], "c") + assert settings["key"] == "c" settings["key2"] = "x" - self.assertIn("key2", settings) - self.assertEqual(settings["key2"], "x") - self.assertEqual(settings.getpriority("key2"), 20) + assert "key2" in settings + assert settings["key2"] == "x" + assert settings.getpriority("key2") == 20 def test_setdict_alias(self): with mock.patch.object(self.settings, "set") as mock_set: self.settings.setdict({"TEST_1": "value1", "TEST_2": "value2"}, 10) - self.assertEqual(mock_set.call_count, 2) + assert mock_set.call_count == 2 calls = [ mock.call("TEST_1", "value1", 10), mock.call("TEST_2", "value2", 10), @@ -149,10 +149,10 @@ class BaseSettingsTest(unittest.TestCase): self.settings.attributes = {} self.settings.setmodule(ModuleMock(), 10) - self.assertIn("UPPERCASE_VAR", self.settings.attributes) - self.assertNotIn("MIXEDcase_VAR", self.settings.attributes) - self.assertNotIn("lowercase_var", self.settings.attributes) - self.assertEqual(len(self.settings.attributes), 1) + assert "UPPERCASE_VAR" in self.settings.attributes + assert "MIXEDcase_VAR" not in self.settings.attributes + assert "lowercase_var" not in self.settings.attributes + assert len(self.settings.attributes) == 1 def test_setmodule_alias(self): with mock.patch.object(self.settings, "set") as mock_set: @@ -168,13 +168,13 @@ class BaseSettingsTest(unittest.TestCase): self.settings.attributes = {} self.settings.setmodule("tests.test_settings.default_settings", 10) - self.assertCountEqual(self.settings.attributes.keys(), ctrl_attributes.keys()) + assert set(self.settings.attributes) == set(ctrl_attributes) for key in ctrl_attributes: attr = self.settings.attributes[key] ctrl_attr = ctrl_attributes[key] - self.assertEqual(attr.value, ctrl_attr.value) - self.assertEqual(attr.priority, ctrl_attr.priority) + assert attr.value == ctrl_attr.value + assert attr.priority == ctrl_attr.priority def test_update(self): settings = BaseSettings({"key_lowprio": 0}, priority=0) @@ -186,21 +186,21 @@ class BaseSettingsTest(unittest.TestCase): custom_dict = {"key_lowprio": 2, "key_highprio": 12, "newkey_two": None} settings.update(custom_dict, priority=20) - self.assertEqual(settings["key_lowprio"], 2) - self.assertEqual(settings.getpriority("key_lowprio"), 20) - self.assertEqual(settings["key_highprio"], 10) - self.assertIn("newkey_two", settings) - self.assertEqual(settings.getpriority("newkey_two"), 20) + assert settings["key_lowprio"] == 2 + assert settings.getpriority("key_lowprio") == 20 + assert settings["key_highprio"] == 10 + assert "newkey_two" in settings + assert settings.getpriority("newkey_two") == 20 settings.update(custom_settings) - self.assertEqual(settings["key_lowprio"], 1) - self.assertEqual(settings.getpriority("key_lowprio"), 30) - self.assertEqual(settings["key_highprio"], 10) - self.assertIn("newkey_one", settings) - self.assertEqual(settings.getpriority("newkey_one"), 50) + assert settings["key_lowprio"] == 1 + assert settings.getpriority("key_lowprio") == 30 + assert settings["key_highprio"] == 10 + assert "newkey_one" in settings + assert settings.getpriority("newkey_one") == 50 settings.update({"key_lowprio": 3}, priority=20) - self.assertEqual(settings["key_lowprio"], 1) + assert settings["key_lowprio"] == 1 @pytest.mark.xfail( raises=TypeError, reason="BaseSettings.update doesn't support kwargs input" @@ -220,21 +220,21 @@ class BaseSettingsTest(unittest.TestCase): def test_update_jsonstring(self): settings = BaseSettings({"number": 0, "dict": BaseSettings({"key": "val"})}) settings.update('{"number": 1, "newnumber": 2}') - self.assertEqual(settings["number"], 1) - self.assertEqual(settings["newnumber"], 2) + assert settings["number"] == 1 + assert settings["newnumber"] == 2 settings.set("dict", '{"key": "newval", "newkey": "newval2"}') - self.assertEqual(settings["dict"]["key"], "newval") - self.assertEqual(settings["dict"]["newkey"], "newval2") + assert settings["dict"]["key"] == "newval" + assert settings["dict"]["newkey"] == "newval2" def test_delete(self): settings = BaseSettings({"key": None}) settings.set("key_highprio", None, priority=50) settings.delete("key") settings.delete("key_highprio") - self.assertNotIn("key", settings) - self.assertIn("key_highprio", settings) + assert "key" not in settings + assert "key_highprio" in settings del settings["key_highprio"] - self.assertNotIn("key_highprio", settings) + assert "key_highprio" not in settings with pytest.raises(KeyError): settings.delete("notkey") with pytest.raises(KeyError): @@ -271,40 +271,40 @@ class BaseSettingsTest(unittest.TestCase): for key, value in test_configuration.items() } - self.assertTrue(settings.getbool("TEST_ENABLED1")) - self.assertTrue(settings.getbool("TEST_ENABLED2")) - self.assertTrue(settings.getbool("TEST_ENABLED3")) - self.assertTrue(settings.getbool("TEST_ENABLED4")) - self.assertTrue(settings.getbool("TEST_ENABLED5")) - self.assertFalse(settings.getbool("TEST_ENABLEDx")) - self.assertTrue(settings.getbool("TEST_ENABLEDx", True)) - self.assertFalse(settings.getbool("TEST_DISABLED1")) - self.assertFalse(settings.getbool("TEST_DISABLED2")) - self.assertFalse(settings.getbool("TEST_DISABLED3")) - self.assertFalse(settings.getbool("TEST_DISABLED4")) - self.assertFalse(settings.getbool("TEST_DISABLED5")) - self.assertEqual(settings.getint("TEST_INT1"), 123) - self.assertEqual(settings.getint("TEST_INT2"), 123) - self.assertEqual(settings.getint("TEST_INTx"), 0) - self.assertEqual(settings.getint("TEST_INTx", 45), 45) - self.assertEqual(settings.getfloat("TEST_FLOAT1"), 123.45) - self.assertEqual(settings.getfloat("TEST_FLOAT2"), 123.45) - self.assertEqual(settings.getfloat("TEST_FLOATx"), 0.0) - self.assertEqual(settings.getfloat("TEST_FLOATx", 55.0), 55.0) - self.assertEqual(settings.getlist("TEST_LIST1"), ["one", "two"]) - self.assertEqual(settings.getlist("TEST_LIST2"), ["one", "two"]) - self.assertEqual(settings.getlist("TEST_LIST3"), []) - self.assertEqual(settings.getlist("TEST_LISTx"), []) - self.assertEqual(settings.getlist("TEST_LISTx", ["default"]), ["default"]) - self.assertEqual(settings["TEST_STR"], "value") - self.assertEqual(settings.get("TEST_STR"), "value") - self.assertEqual(settings["TEST_STRx"], None) - self.assertEqual(settings.get("TEST_STRx"), None) - self.assertEqual(settings.get("TEST_STRx", "default"), "default") - self.assertEqual(settings.getdict("TEST_DICT1"), {"key1": "val1", "ke2": 3}) - self.assertEqual(settings.getdict("TEST_DICT2"), {"key1": "val1", "ke2": 3}) - self.assertEqual(settings.getdict("TEST_DICT3"), {}) - self.assertEqual(settings.getdict("TEST_DICT3", {"key1": 5}), {"key1": 5}) + assert settings.getbool("TEST_ENABLED1") + assert settings.getbool("TEST_ENABLED2") + assert settings.getbool("TEST_ENABLED3") + assert settings.getbool("TEST_ENABLED4") + assert settings.getbool("TEST_ENABLED5") + assert not settings.getbool("TEST_ENABLEDx") + assert settings.getbool("TEST_ENABLEDx", True) + assert not settings.getbool("TEST_DISABLED1") + assert not settings.getbool("TEST_DISABLED2") + assert not settings.getbool("TEST_DISABLED3") + assert not settings.getbool("TEST_DISABLED4") + assert not settings.getbool("TEST_DISABLED5") + assert settings.getint("TEST_INT1") == 123 + assert settings.getint("TEST_INT2") == 123 + assert settings.getint("TEST_INTx") == 0 + assert settings.getint("TEST_INTx", 45) == 45 + assert settings.getfloat("TEST_FLOAT1") == 123.45 + assert settings.getfloat("TEST_FLOAT2") == 123.45 + assert settings.getfloat("TEST_FLOATx") == 0.0 + assert settings.getfloat("TEST_FLOATx", 55.0) == 55.0 + assert settings.getlist("TEST_LIST1") == ["one", "two"] + assert settings.getlist("TEST_LIST2") == ["one", "two"] + assert settings.getlist("TEST_LIST3") == [] + assert settings.getlist("TEST_LISTx") == [] + assert settings.getlist("TEST_LISTx", ["default"]) == ["default"] + assert settings["TEST_STR"] == "value" + assert settings.get("TEST_STR") == "value" + assert settings["TEST_STRx"] is None + assert settings.get("TEST_STRx") is None + assert settings.get("TEST_STRx", "default") == "default" + assert settings.getdict("TEST_DICT1") == {"key1": "val1", "ke2": 3} + assert settings.getdict("TEST_DICT2") == {"key1": "val1", "ke2": 3} + assert settings.getdict("TEST_DICT3") == {} + assert settings.getdict("TEST_DICT3", {"key1": 5}) == {"key1": 5} with pytest.raises( ValueError, match="dictionary update sequence element #0 has length 3; 2 is required|sequence of pairs expected", @@ -321,8 +321,8 @@ class BaseSettingsTest(unittest.TestCase): def test_getpriority(self): settings = BaseSettings({"key": "value"}, priority=99) - self.assertEqual(settings.getpriority("key"), 99) - self.assertEqual(settings.getpriority("nonexistentkey"), None) + assert settings.getpriority("key") == 99 + assert settings.getpriority("nonexistentkey") is None def test_getwithbase(self): s = BaseSettings( @@ -333,16 +333,16 @@ class BaseSettingsTest(unittest.TestCase): } ) s["TEST"].set(2, 200, "cmdline") - self.assertCountEqual(s.getwithbase("TEST"), {1: 1, 2: 200, 3: 30}) - self.assertCountEqual(s.getwithbase("HASNOBASE"), s["HASNOBASE"]) - self.assertEqual(s.getwithbase("NONEXISTENT"), {}) + assert set(s.getwithbase("TEST")) == {1, 2, 3} + assert set(s.getwithbase("HASNOBASE")) == set(s["HASNOBASE"]) + assert s.getwithbase("NONEXISTENT") == {} def test_maxpriority(self): # Empty settings should return 'default' - self.assertEqual(self.settings.maxpriority(), 0) + assert self.settings.maxpriority() == 0 self.settings.set("A", 0, 10) self.settings.set("B", 0, 30) - self.assertEqual(self.settings.maxpriority(), 30) + assert self.settings.maxpriority() == 30 def test_copy(self): values = { @@ -356,17 +356,15 @@ class BaseSettingsTest(unittest.TestCase): self.settings.setdict(values) copy = self.settings.copy() self.settings.set("TEST_BOOL", False) - self.assertTrue(copy.get("TEST_BOOL")) + assert copy.get("TEST_BOOL") test_list = self.settings.get("TEST_LIST") test_list.append("three") - self.assertListEqual(copy.get("TEST_LIST"), ["one", "two"]) + assert copy.get("TEST_LIST") == ["one", "two"] test_list_of_lists = self.settings.get("TEST_LIST_OF_LISTS") test_list_of_lists[0].append("first_three") - self.assertListEqual( - copy.get("TEST_LIST_OF_LISTS")[0], ["first_one", "first_two"] - ) + assert copy.get("TEST_LIST_OF_LISTS")[0] == ["first_one", "first_two"] def test_copy_to_dict(self): s = BaseSettings( @@ -379,17 +377,14 @@ class BaseSettingsTest(unittest.TestCase): "HASNOBASE": BaseSettings({3: 3000}, "default"), } ) - self.assertDictEqual( - s.copy_to_dict(), - { - "HASNOBASE": {3: 3000}, - "TEST": {1: 10, 3: 30}, - "TEST_BASE": {1: 1, 2: 2}, - "TEST_LIST": [1, 2], - "TEST_BOOLEAN": False, - "TEST_STRING": "a string", - }, - ) + assert s.copy_to_dict() == { + "HASNOBASE": {3: 3000}, + "TEST": {1: 10, 3: 30}, + "TEST_BASE": {1: 1, 2: 2}, + "TEST_LIST": [1, 2], + "TEST_BOOLEAN": False, + "TEST_STRING": "a string", + } def test_freeze(self): self.settings.freeze() @@ -400,55 +395,55 @@ class BaseSettingsTest(unittest.TestCase): def test_frozencopy(self): frozencopy = self.settings.frozencopy() - self.assertTrue(frozencopy.frozen) - self.assertIsNot(frozencopy, self.settings) + assert frozencopy.frozen + assert frozencopy is not self.settings -class SettingsTest(unittest.TestCase): - def setUp(self): +class TestSettings: + def setup_method(self): self.settings = Settings() @mock.patch.dict("scrapy.settings.SETTINGS_PRIORITIES", {"default": 10}) @mock.patch("scrapy.settings.default_settings", default_settings) def test_initial_defaults(self): settings = Settings() - self.assertEqual(len(settings.attributes), 2) - self.assertIn("TEST_DEFAULT", settings.attributes) + assert len(settings.attributes) == 2 + assert "TEST_DEFAULT" in settings.attributes attr = settings.attributes["TEST_DEFAULT"] - self.assertIsInstance(attr, SettingsAttribute) - self.assertEqual(attr.value, "defvalue") - self.assertEqual(attr.priority, 10) + assert isinstance(attr, SettingsAttribute) + assert attr.value == "defvalue" + assert attr.priority == 10 @mock.patch.dict("scrapy.settings.SETTINGS_PRIORITIES", {}) @mock.patch("scrapy.settings.default_settings", {}) def test_initial_values(self): settings = Settings({"TEST_OPTION": "value"}, 10) - self.assertEqual(len(settings.attributes), 1) - self.assertIn("TEST_OPTION", settings.attributes) + assert len(settings.attributes) == 1 + assert "TEST_OPTION" in settings.attributes attr = settings.attributes["TEST_OPTION"] - self.assertIsInstance(attr, SettingsAttribute) - self.assertEqual(attr.value, "value") - self.assertEqual(attr.priority, 10) + assert isinstance(attr, SettingsAttribute) + assert attr.value == "value" + assert attr.priority == 10 @mock.patch("scrapy.settings.default_settings", default_settings) def test_autopromote_dicts(self): settings = Settings() mydict = settings.get("TEST_DICT") - self.assertIsInstance(mydict, BaseSettings) - self.assertIn("key", mydict) - self.assertEqual(mydict["key"], "val") # pylint: disable=unsubscriptable-object - self.assertEqual(mydict.getpriority("key"), 0) + assert isinstance(mydict, BaseSettings) + assert "key" in mydict + assert mydict["key"] == "val" + assert mydict.getpriority("key") == 0 @mock.patch("scrapy.settings.default_settings", default_settings) def test_getdict_autodegrade_basesettings(self): settings = Settings() mydict = settings.getdict("TEST_DICT") - self.assertIsInstance(mydict, dict) - self.assertEqual(len(mydict), 1) - self.assertIn("key", mydict) - self.assertEqual(mydict["key"], "val") + assert isinstance(mydict, dict) + assert len(mydict) == 1 + assert "key" in mydict + assert mydict["key"] == "val" def test_passing_objects_as_values(self): from scrapy.core.downloader.handlers.file import FileDownloadHandler @@ -470,19 +465,19 @@ class SettingsTest(unittest.TestCase): } ) - self.assertIn("ITEM_PIPELINES", settings.attributes) + assert "ITEM_PIPELINES" in settings.attributes mypipeline, priority = settings.getdict("ITEM_PIPELINES").popitem() - self.assertEqual(priority, 800) - self.assertEqual(mypipeline, TestPipeline) - self.assertIsInstance(mypipeline(), TestPipeline) - self.assertEqual(mypipeline().process_item("item", None), "item") + assert priority == 800 + assert mypipeline == TestPipeline + assert isinstance(mypipeline(), TestPipeline) + assert mypipeline().process_item("item", None) == "item" myhandler = settings.getdict("DOWNLOAD_HANDLERS").pop("ftp") - self.assertEqual(myhandler, FileDownloadHandler) + assert myhandler == FileDownloadHandler myhandler_instance = build_from_crawler(myhandler, get_crawler()) - self.assertIsInstance(myhandler_instance, FileDownloadHandler) - self.assertTrue(hasattr(myhandler_instance, "download_request")) + assert isinstance(myhandler_instance, FileDownloadHandler) + assert hasattr(myhandler_instance, "download_request") def test_pop_item_with_default_value(self): settings = Settings() @@ -491,14 +486,14 @@ class SettingsTest(unittest.TestCase): settings.pop("DUMMY_CONFIG") dummy_config_value = settings.pop("DUMMY_CONFIG", "dummy_value") - self.assertEqual(dummy_config_value, "dummy_value") + assert dummy_config_value == "dummy_value" def test_pop_item_with_immutable_settings(self): settings = Settings( {"DUMMY_CONFIG": "dummy_value", "OTHER_DUMMY_CONFIG": "other_dummy_value"} ) - self.assertEqual(settings.pop("DUMMY_CONFIG"), "dummy_value") + assert settings.pop("DUMMY_CONFIG") == "dummy_value" settings.freeze() From 7bbe775040d5d695bc3b48a73de5d6fa99312b4c Mon Sep 17 00:00:00 2001 From: Andrey Rakhmatullin Date: Sun, 9 Mar 2025 23:24:45 +0400 Subject: [PATCH 5/5] Converting tests to plain asserts, part 5. (#6712) --- tests/test_logformatter.py | 93 +++++++-------- tests/test_logstats.py | 45 ++++---- tests/test_mail.py | 71 ++++++------ tests/test_middleware.py | 30 ++--- tests/test_pipelines.py | 16 +-- tests/test_pqueues.py | 113 +++++++++--------- tests/test_proxy_connect.py | 11 +- tests/test_request_attribute_binding.py | 26 ++--- tests/test_request_cb_kwargs.py | 30 ++--- tests/test_request_dict.py | 50 ++++---- tests/test_request_left.py | 8 +- tests/test_responsetypes.py | 8 +- tests/test_robotstxt_interface.py | 106 +++++++---------- tests/test_scheduler.py | 119 +++++++++---------- tests/test_scheduler_base.py | 59 +++++----- tests/test_selector.py | 145 +++++++++++------------- tests/test_signals.py | 6 +- tests/test_toplevel.py | 20 ++-- tests/test_urlparse_monkeypatches.py | 11 +- 19 files changed, 446 insertions(+), 521 deletions(-) diff --git a/tests/test_logformatter.py b/tests/test_logformatter.py index 962692a31..3c9f97631 100644 --- a/tests/test_logformatter.py +++ b/tests/test_logformatter.py @@ -1,11 +1,10 @@ import logging -import unittest import pytest from testfixtures import LogCapture from twisted.internet import defer from twisted.python.failure import Failure -from twisted.trial.unittest import TestCase as TwistedTestCase +from twisted.trial.unittest import TestCase from scrapy.exceptions import DropItem from scrapy.http import Request, Response @@ -24,8 +23,8 @@ class CustomItem(Item): return f"name: {self['name']}" -class LogFormatterTestCase(unittest.TestCase): - def setUp(self): +class TestLogFormatter: + def setup_method(self): self.formatter = LogFormatter() self.spider = Spider("default") self.spider.crawler = get_crawler() @@ -35,9 +34,7 @@ class LogFormatterTestCase(unittest.TestCase): res = Response("http://www.example.com") logkws = self.formatter.crawled(req, res, self.spider) logline = logkws["msg"] % logkws["args"] - self.assertEqual( - logline, "Crawled (200) (referer: None)" - ) + assert logline == "Crawled (200) (referer: None)" def test_crawled_without_referer(self): req = Request( @@ -46,9 +43,9 @@ class LogFormatterTestCase(unittest.TestCase): res = Response("http://www.example.com", flags=["cached"]) logkws = self.formatter.crawled(req, res, self.spider) logline = logkws["msg"] % logkws["args"] - self.assertEqual( - logline, - "Crawled (200) (referer: http://example.com) ['cached']", + assert ( + logline + == "Crawled (200) (referer: http://example.com) ['cached']" ) def test_flags_in_request(self): @@ -56,9 +53,9 @@ class LogFormatterTestCase(unittest.TestCase): res = Response("http://www.example.com") logkws = self.formatter.crawled(req, res, self.spider) logline = logkws["msg"] % logkws["args"] - self.assertEqual( - logline, - "Crawled (200) ['test', 'flag'] (referer: None)", + assert ( + logline + == "Crawled (200) ['test', 'flag'] (referer: None)" ) def test_dropped(self): @@ -69,7 +66,7 @@ class LogFormatterTestCase(unittest.TestCase): logline = logkws["msg"] % logkws["args"] lines = logline.splitlines() assert all(isinstance(x, str) for x in lines) - self.assertEqual(lines, ["Dropped: \u2018", "{}"]) + assert lines == ["Dropped: \u2018", "{}"] def test_dropitem_default_log_level(self): item = {} @@ -79,38 +76,38 @@ class LogFormatterTestCase(unittest.TestCase): spider.crawler = get_crawler(Spider) logkws = self.formatter.dropped(item, exception, response, spider) - self.assertEqual(logkws["level"], logging.WARNING) + assert logkws["level"] == logging.WARNING spider.crawler.settings.frozen = False spider.crawler.settings["DEFAULT_DROPITEM_LOG_LEVEL"] = logging.INFO spider.crawler.settings.frozen = True logkws = self.formatter.dropped(item, exception, response, spider) - self.assertEqual(logkws["level"], logging.INFO) + assert logkws["level"] == logging.INFO spider.crawler.settings.frozen = False spider.crawler.settings["DEFAULT_DROPITEM_LOG_LEVEL"] = "INFO" spider.crawler.settings.frozen = True logkws = self.formatter.dropped(item, exception, response, spider) - self.assertEqual(logkws["level"], logging.INFO) + assert logkws["level"] == logging.INFO spider.crawler.settings.frozen = False spider.crawler.settings["DEFAULT_DROPITEM_LOG_LEVEL"] = 10 spider.crawler.settings.frozen = True logkws = self.formatter.dropped(item, exception, response, spider) - self.assertEqual(logkws["level"], logging.DEBUG) + assert logkws["level"] == logging.DEBUG spider.crawler.settings.frozen = False spider.crawler.settings["DEFAULT_DROPITEM_LOG_LEVEL"] = 0 spider.crawler.settings.frozen = True logkws = self.formatter.dropped(item, exception, response, spider) - self.assertEqual(logkws["level"], logging.NOTSET) + assert logkws["level"] == logging.NOTSET unsupported_value = object() spider.crawler.settings.frozen = False spider.crawler.settings["DEFAULT_DROPITEM_LOG_LEVEL"] = unsupported_value spider.crawler.settings.frozen = True logkws = self.formatter.dropped(item, exception, response, spider) - self.assertEqual(logkws["level"], unsupported_value) + assert logkws["level"] == unsupported_value with pytest.raises(TypeError): logging.log(logkws["level"], "message") @@ -121,11 +118,11 @@ class LogFormatterTestCase(unittest.TestCase): exception = DropItem("Test drop", log_level="INFO") logkws = self.formatter.dropped(item, exception, response, self.spider) - self.assertEqual(logkws["level"], logging.INFO) + assert logkws["level"] == logging.INFO exception = DropItem("Test drop", log_level="ERROR") logkws = self.formatter.dropped(item, exception, response, self.spider) - self.assertEqual(logkws["level"], logging.ERROR) + assert logkws["level"] == logging.ERROR def test_item_error(self): # In practice, the complete traceback is shown by passing the @@ -135,7 +132,7 @@ class LogFormatterTestCase(unittest.TestCase): response = Response("http://www.example.com") logkws = self.formatter.item_error(item, exception, response, self.spider) logline = logkws["msg"] % logkws["args"] - self.assertEqual(logline, "Error processing {'key': 'value'}") + assert logline == "Error processing {'key': 'value'}" def test_spider_error(self): # In practice, the complete traceback is shown by passing the @@ -147,9 +144,9 @@ class LogFormatterTestCase(unittest.TestCase): response = Response("http://www.example.com", request=request) logkws = self.formatter.spider_error(failure, request, response, self.spider) logline = logkws["msg"] % logkws["args"] - self.assertEqual( - logline, - "Spider error processing (referer: http://example.org)", + assert ( + logline + == "Spider error processing (referer: http://example.org)" ) def test_download_error_short(self): @@ -159,7 +156,7 @@ class LogFormatterTestCase(unittest.TestCase): request = Request("http://www.example.com") logkws = self.formatter.download_error(failure, request, self.spider) logline = logkws["msg"] % logkws["args"] - self.assertEqual(logline, "Error downloading ") + assert logline == "Error downloading " def test_download_error_long(self): # In practice, the complete traceback is shown by passing the @@ -170,9 +167,7 @@ class LogFormatterTestCase(unittest.TestCase): failure, request, self.spider, "Some message" ) logline = logkws["msg"] % logkws["args"] - self.assertEqual( - logline, "Error downloading : Some message" - ) + assert logline == "Error downloading : Some message" def test_scraped(self): item = CustomItem() @@ -182,9 +177,7 @@ class LogFormatterTestCase(unittest.TestCase): logline = logkws["msg"] % logkws["args"] lines = logline.splitlines() assert all(isinstance(x, str) for x in lines) - self.assertEqual( - lines, ["Scraped from <200 http://www.example.com>", "name: \xa3"] - ) + assert lines == ["Scraped from <200 http://www.example.com>", "name: \xa3"] class LogFormatterSubclass(LogFormatter): @@ -200,8 +193,8 @@ class LogFormatterSubclass(LogFormatter): } -class LogformatterSubclassTest(LogFormatterTestCase): - def setUp(self): +class TestLogformatterSubclass(TestLogFormatter): + def setup_method(self): self.formatter = LogFormatterSubclass() self.spider = Spider("default") self.spider.crawler = get_crawler(Spider) @@ -211,8 +204,8 @@ class LogformatterSubclassTest(LogFormatterTestCase): res = Response("http://www.example.com") logkws = self.formatter.crawled(req, res, self.spider) logline = logkws["msg"] % logkws["args"] - self.assertEqual( - logline, "Crawled (200) (referer: None) []" + assert ( + logline == "Crawled (200) (referer: None) []" ) def test_crawled_without_referer(self): @@ -224,9 +217,9 @@ class LogformatterSubclassTest(LogFormatterTestCase): res = Response("http://www.example.com") logkws = self.formatter.crawled(req, res, self.spider) logline = logkws["msg"] % logkws["args"] - self.assertEqual( - logline, - "Crawled (200) (referer: http://example.com) ['cached']", + assert ( + logline + == "Crawled (200) (referer: http://example.com) ['cached']" ) def test_flags_in_request(self): @@ -234,9 +227,9 @@ class LogformatterSubclassTest(LogFormatterTestCase): res = Response("http://www.example.com") logkws = self.formatter.crawled(req, res, self.spider) logline = logkws["msg"] % logkws["args"] - self.assertEqual( - logline, - "Crawled (200) (referer: None) ['test', 'flag']", + assert ( + logline + == "Crawled (200) (referer: None) ['test', 'flag']" ) @@ -261,7 +254,7 @@ class DropSomeItemsPipeline: self.drop = True -class ShowOrSkipMessagesTestCase(TwistedTestCase): +class TestShowOrSkipMessages(TestCase): @classmethod def setUpClass(cls): cls.mockserver = MockServer() @@ -284,9 +277,9 @@ class ShowOrSkipMessagesTestCase(TwistedTestCase): crawler = get_crawler(ItemSpider, self.base_settings) with LogCapture() as lc: yield crawler.crawl(mockserver=self.mockserver) - self.assertIn("Scraped from <200 http://127.0.0.1:", str(lc)) - self.assertIn("Crawled (200) body

") - self.assertEqual(msg.get("Content-Type"), "text/html") + assert msg.get_payload() == "

body

" + assert msg.get("Content-Type") == "text/html" def test_send_attach(self): attach = BytesIO() @@ -70,22 +69,22 @@ class MailSenderTest(unittest.TestCase): ) assert self.catched_msg - self.assertEqual(self.catched_msg["to"], ["test@scrapy.org"]) - self.assertEqual(self.catched_msg["subject"], "subject") - self.assertEqual(self.catched_msg["body"], "body") + assert self.catched_msg["to"] == ["test@scrapy.org"] + assert self.catched_msg["subject"] == "subject" + assert self.catched_msg["body"] == "body" msg = self.catched_msg["msg"] - self.assertEqual(msg["to"], "test@scrapy.org") - self.assertEqual(msg["subject"], "subject") + assert msg["to"] == "test@scrapy.org" + assert msg["subject"] == "subject" payload = msg.get_payload() assert isinstance(payload, list) - self.assertEqual(len(payload), 2) + assert len(payload) == 2 text, attach = payload - self.assertEqual(text.get_payload(decode=True), b"body") - self.assertEqual(text.get_charset(), Charset("us-ascii")) - self.assertEqual(attach.get_payload(decode=True), b"content") + assert text.get_payload(decode=True) == b"body" + assert text.get_charset() == Charset("us-ascii") + assert attach.get_payload(decode=True) == b"content" def _catch_mail_sent(self, **kwargs): self.catched_msg = {**kwargs} @@ -103,14 +102,14 @@ class MailSenderTest(unittest.TestCase): ) assert self.catched_msg - self.assertEqual(self.catched_msg["subject"], subject) - self.assertEqual(self.catched_msg["body"], body) + assert self.catched_msg["subject"] == subject + assert self.catched_msg["body"] == body msg = self.catched_msg["msg"] - self.assertEqual(msg["subject"], subject) - self.assertEqual(msg.get_payload(decode=True).decode("utf-8"), body) - self.assertEqual(msg.get_charset(), Charset("utf-8")) - self.assertEqual(msg.get("Content-Type"), 'text/plain; charset="utf-8"') + assert msg["subject"] == subject + assert msg.get_payload(decode=True).decode("utf-8") == body + assert msg.get_charset() == Charset("utf-8") + assert msg.get("Content-Type") == 'text/plain; charset="utf-8"' def test_send_attach_utf8(self): subject = "sübjèçt" @@ -131,22 +130,22 @@ class MailSenderTest(unittest.TestCase): ) assert self.catched_msg - self.assertEqual(self.catched_msg["subject"], subject) - self.assertEqual(self.catched_msg["body"], body) + assert self.catched_msg["subject"] == subject + assert self.catched_msg["body"] == body msg = self.catched_msg["msg"] - self.assertEqual(msg["subject"], subject) - self.assertEqual(msg.get_charset(), Charset("utf-8")) - self.assertEqual(msg.get("Content-Type"), 'multipart/mixed; charset="utf-8"') + assert msg["subject"] == subject + assert msg.get_charset() == Charset("utf-8") + assert msg.get("Content-Type") == 'multipart/mixed; charset="utf-8"' payload = msg.get_payload() assert isinstance(payload, list) - self.assertEqual(len(payload), 2) + assert len(payload) == 2 text, attach = payload - self.assertEqual(text.get_payload(decode=True).decode("utf-8"), body) - self.assertEqual(text.get_charset(), Charset("utf-8")) - self.assertEqual(attach.get_payload(decode=True).decode("utf-8"), body) + assert text.get_payload(decode=True).decode("utf-8") == body + assert text.get_charset() == Charset("utf-8") + assert attach.get_payload(decode=True).decode("utf-8") == body def test_create_sender_factory_with_host(self): mailsender = MailSender(debug=False, smtphost="smtp.testhost.com") @@ -156,4 +155,4 @@ class MailSenderTest(unittest.TestCase): ) context = factory.buildProtocol("test@scrapy.org").context - self.assertIsInstance(context, ClientTLSOptions) + assert isinstance(context, ClientTLSOptions) diff --git a/tests/test_middleware.py b/tests/test_middleware.py index 0cc532570..d004d4d93 100644 --- a/tests/test_middleware.py +++ b/tests/test_middleware.py @@ -1,5 +1,3 @@ -from twisted.trial import unittest - from scrapy.exceptions import NotConfigured from scrapy.middleware import MiddlewareManager from scrapy.utils.test import get_crawler @@ -51,37 +49,27 @@ class MyMiddlewareManager(MiddlewareManager): self.methods["process"].append(mw.process) -class MiddlewareManagerTest(unittest.TestCase): +class TestMiddlewareManager: def test_init(self): m1, m2, m3 = M1(), M2(), M3() mwman = MyMiddlewareManager(m1, m2, m3) - self.assertEqual( - list(mwman.methods["open_spider"]), [m1.open_spider, m2.open_spider] - ) - self.assertEqual( - list(mwman.methods["close_spider"]), [m2.close_spider, m1.close_spider] - ) - self.assertEqual(list(mwman.methods["process"]), [m1.process, m3.process]) + assert list(mwman.methods["open_spider"]) == [m1.open_spider, m2.open_spider] + assert list(mwman.methods["close_spider"]) == [m2.close_spider, m1.close_spider] + assert list(mwman.methods["process"]) == [m1.process, m3.process] def test_methods(self): mwman = MyMiddlewareManager(M1(), M2(), M3()) - self.assertEqual( - [x.__self__.__class__ for x in mwman.methods["open_spider"]], [M1, M2] - ) - self.assertEqual( - [x.__self__.__class__ for x in mwman.methods["close_spider"]], [M2, M1] - ) - self.assertEqual( - [x.__self__.__class__ for x in mwman.methods["process"]], [M1, M3] - ) + assert [x.__self__.__class__ for x in mwman.methods["open_spider"]] == [M1, M2] + assert [x.__self__.__class__ for x in mwman.methods["close_spider"]] == [M2, M1] + assert [x.__self__.__class__ for x in mwman.methods["process"]] == [M1, M3] def test_enabled(self): m1, m2, m3 = M1(), M2(), M3() mwman = MiddlewareManager(m1, m2, m3) - self.assertEqual(mwman.middlewares, (m1, m2, m3)) + assert mwman.middlewares == (m1, m2, m3) def test_enabled_from_settings(self): crawler = get_crawler() mwman = MyMiddlewareManager.from_crawler(crawler) classes = [x.__class__ for x in mwman.middlewares] - self.assertEqual(classes, [M1, M3]) + assert classes == [M1, M3] diff --git a/tests/test_pipelines.py b/tests/test_pipelines.py index 0ae86235c..743d9774b 100644 --- a/tests/test_pipelines.py +++ b/tests/test_pipelines.py @@ -76,7 +76,7 @@ class ItemSpider(Spider): return {"field": 42} -class PipelineTestCase(unittest.TestCase): +class TestPipeline(unittest.TestCase): @classmethod def setUpClass(cls): cls.mockserver = MockServer() @@ -87,8 +87,8 @@ class PipelineTestCase(unittest.TestCase): cls.mockserver.__exit__(None, None, None) def _on_item_scraped(self, item): - self.assertIsInstance(item, dict) - self.assertTrue(item.get("pipeline_passed")) + assert isinstance(item, dict) + assert item.get("pipeline_passed") self.items.append(item) def _create_crawler(self, pipeline_class): @@ -104,30 +104,30 @@ class PipelineTestCase(unittest.TestCase): def test_simple_pipeline(self): crawler = self._create_crawler(SimplePipeline) yield crawler.crawl(mockserver=self.mockserver) - self.assertEqual(len(self.items), 1) + assert len(self.items) == 1 @defer.inlineCallbacks def test_deferred_pipeline(self): crawler = self._create_crawler(DeferredPipeline) yield crawler.crawl(mockserver=self.mockserver) - self.assertEqual(len(self.items), 1) + assert len(self.items) == 1 @defer.inlineCallbacks def test_asyncdef_pipeline(self): crawler = self._create_crawler(AsyncDefPipeline) yield crawler.crawl(mockserver=self.mockserver) - self.assertEqual(len(self.items), 1) + assert len(self.items) == 1 @pytest.mark.only_asyncio @defer.inlineCallbacks def test_asyncdef_asyncio_pipeline(self): crawler = self._create_crawler(AsyncDefAsyncioPipeline) yield crawler.crawl(mockserver=self.mockserver) - self.assertEqual(len(self.items), 1) + assert len(self.items) == 1 @pytest.mark.only_not_asyncio @defer.inlineCallbacks def test_asyncdef_not_asyncio_pipeline(self): crawler = self._create_crawler(AsyncDefNotAsyncioPipeline) yield crawler.crawl(mockserver=self.mockserver) - self.assertEqual(len(self.items), 1) + assert len(self.items) == 1 diff --git a/tests/test_pqueues.py b/tests/test_pqueues.py index c223c4562..d5c710ed2 100644 --- a/tests/test_pqueues.py +++ b/tests/test_pqueues.py @@ -1,5 +1,4 @@ import tempfile -import unittest import pytest import queuelib @@ -12,8 +11,8 @@ from scrapy.utils.test import get_crawler from tests.test_scheduler import MockDownloader, MockEngine -class PriorityQueueTest(unittest.TestCase): - def setUp(self): +class TestPriorityQueue: + def setup_method(self): self.crawler = get_crawler(Spider) self.spider = self.crawler._create_spider("foo") @@ -22,20 +21,20 @@ class PriorityQueueTest(unittest.TestCase): queue = ScrapyPriorityQueue.from_crawler( self.crawler, FifoMemoryQueue, temp_dir ) - self.assertIsNone(queue.pop()) - self.assertEqual(len(queue), 0) + assert queue.pop() is None + assert len(queue) == 0 req1 = Request("https://example.org/1", priority=1) queue.push(req1) - self.assertEqual(len(queue), 1) + assert len(queue) == 1 dequeued = queue.pop() - self.assertEqual(len(queue), 0) - self.assertEqual(dequeued.url, req1.url) - self.assertEqual(dequeued.priority, req1.priority) - self.assertEqual(queue.close(), []) + assert len(queue) == 0 + assert dequeued.url == req1.url + assert dequeued.priority == req1.priority + assert not queue.close() def test_no_peek_raises(self): if hasattr(queuelib.queue.FifoMemoryQueue, "peek"): - raise unittest.SkipTest("queuelib.queue.FifoMemoryQueue.peek is defined") + pytest.skip("queuelib.queue.FifoMemoryQueue.peek is defined") temp_dir = tempfile.mkdtemp() queue = ScrapyPriorityQueue.from_crawler( self.crawler, FifoMemoryQueue, temp_dir @@ -50,53 +49,53 @@ class PriorityQueueTest(unittest.TestCase): def test_peek(self): if not hasattr(queuelib.queue.FifoMemoryQueue, "peek"): - raise unittest.SkipTest("queuelib.queue.FifoMemoryQueue.peek is undefined") + pytest.skip("queuelib.queue.FifoMemoryQueue.peek is undefined") temp_dir = tempfile.mkdtemp() queue = ScrapyPriorityQueue.from_crawler( self.crawler, FifoMemoryQueue, temp_dir ) - self.assertEqual(len(queue), 0) - self.assertIsNone(queue.peek()) + assert len(queue) == 0 + assert queue.peek() is None req1 = Request("https://example.org/1") req2 = Request("https://example.org/2") req3 = Request("https://example.org/3") queue.push(req1) queue.push(req2) queue.push(req3) - self.assertEqual(len(queue), 3) - self.assertEqual(queue.peek().url, req1.url) - self.assertEqual(queue.pop().url, req1.url) - self.assertEqual(len(queue), 2) - self.assertEqual(queue.peek().url, req2.url) - self.assertEqual(queue.pop().url, req2.url) - self.assertEqual(len(queue), 1) - self.assertEqual(queue.peek().url, req3.url) - self.assertEqual(queue.pop().url, req3.url) - self.assertEqual(queue.close(), []) + assert len(queue) == 3 + assert queue.peek().url == req1.url + assert queue.pop().url == req1.url + assert len(queue) == 2 + assert queue.peek().url == req2.url + assert queue.pop().url == req2.url + assert len(queue) == 1 + assert queue.peek().url == req3.url + assert queue.pop().url == req3.url + assert not queue.close() def test_queue_push_pop_priorities(self): temp_dir = tempfile.mkdtemp() queue = ScrapyPriorityQueue.from_crawler( self.crawler, FifoMemoryQueue, temp_dir, [-1, -2, -3] ) - self.assertIsNone(queue.pop()) - self.assertEqual(len(queue), 0) + assert queue.pop() is None + assert len(queue) == 0 req1 = Request("https://example.org/1", priority=1) req2 = Request("https://example.org/2", priority=2) req3 = Request("https://example.org/3", priority=3) queue.push(req1) queue.push(req2) queue.push(req3) - self.assertEqual(len(queue), 3) + assert len(queue) == 3 dequeued = queue.pop() - self.assertEqual(len(queue), 2) - self.assertEqual(dequeued.url, req3.url) - self.assertEqual(dequeued.priority, req3.priority) - self.assertEqual(queue.close(), [-1, -2]) + assert len(queue) == 2 + assert dequeued.url == req3.url + assert dequeued.priority == req3.priority + assert queue.close() == [-1, -2] -class DownloaderAwarePriorityQueueTest(unittest.TestCase): - def setUp(self): +class TestDownloaderAwarePriorityQueue: + def setup_method(self): crawler = get_crawler(Spider) crawler.engine = MockEngine(downloader=MockDownloader()) self.queue = DownloaderAwarePriorityQueue.from_crawler( @@ -105,30 +104,30 @@ class DownloaderAwarePriorityQueueTest(unittest.TestCase): key="foo/bar", ) - def tearDown(self): + def teardown_method(self): self.queue.close() def test_push_pop(self): - self.assertEqual(len(self.queue), 0) - self.assertIsNone(self.queue.pop()) + assert len(self.queue) == 0 + assert self.queue.pop() is None req1 = Request("http://www.example.com/1") req2 = Request("http://www.example.com/2") req3 = Request("http://www.example.com/3") self.queue.push(req1) self.queue.push(req2) self.queue.push(req3) - self.assertEqual(len(self.queue), 3) - self.assertEqual(self.queue.pop().url, req1.url) - self.assertEqual(len(self.queue), 2) - self.assertEqual(self.queue.pop().url, req2.url) - self.assertEqual(len(self.queue), 1) - self.assertEqual(self.queue.pop().url, req3.url) - self.assertEqual(len(self.queue), 0) - self.assertIsNone(self.queue.pop()) + assert len(self.queue) == 3 + assert self.queue.pop().url == req1.url + assert len(self.queue) == 2 + assert self.queue.pop().url == req2.url + assert len(self.queue) == 1 + assert self.queue.pop().url == req3.url + assert len(self.queue) == 0 + assert self.queue.pop() is None def test_no_peek_raises(self): if hasattr(queuelib.queue.FifoMemoryQueue, "peek"): - raise unittest.SkipTest("queuelib.queue.FifoMemoryQueue.peek is defined") + pytest.skip("queuelib.queue.FifoMemoryQueue.peek is defined") self.queue.push(Request("https://example.org")) with pytest.raises( NotImplementedError, @@ -138,21 +137,21 @@ class DownloaderAwarePriorityQueueTest(unittest.TestCase): def test_peek(self): if not hasattr(queuelib.queue.FifoMemoryQueue, "peek"): - raise unittest.SkipTest("queuelib.queue.FifoMemoryQueue.peek is undefined") - self.assertEqual(len(self.queue), 0) + pytest.skip("queuelib.queue.FifoMemoryQueue.peek is undefined") + assert len(self.queue) == 0 req1 = Request("https://example.org/1") req2 = Request("https://example.org/2") req3 = Request("https://example.org/3") self.queue.push(req1) self.queue.push(req2) self.queue.push(req3) - self.assertEqual(len(self.queue), 3) - self.assertEqual(self.queue.peek().url, req1.url) - self.assertEqual(self.queue.pop().url, req1.url) - self.assertEqual(len(self.queue), 2) - self.assertEqual(self.queue.peek().url, req2.url) - self.assertEqual(self.queue.pop().url, req2.url) - self.assertEqual(len(self.queue), 1) - self.assertEqual(self.queue.peek().url, req3.url) - self.assertEqual(self.queue.pop().url, req3.url) - self.assertIsNone(self.queue.peek()) + assert len(self.queue) == 3 + assert self.queue.peek().url == req1.url + assert self.queue.pop().url == req1.url + assert len(self.queue) == 2 + assert self.queue.peek().url == req2.url + assert self.queue.pop().url == req2.url + assert len(self.queue) == 1 + assert self.queue.peek().url == req3.url + assert self.queue.pop().url == req3.url + assert self.queue.peek() is None diff --git a/tests/test_proxy_connect.py b/tests/test_proxy_connect.py index 6ed7e93a6..885b7b7ae 100644 --- a/tests/test_proxy_connect.py +++ b/tests/test_proxy_connect.py @@ -6,6 +6,7 @@ from pathlib import Path from subprocess import PIPE, Popen from urllib.parse import urlsplit, urlunsplit +import pytest from testfixtures import LogCapture from twisted.internet import defer from twisted.trial.unittest import TestCase @@ -61,7 +62,7 @@ def _wrong_credentials(proxy_url): return urlunsplit(bad_auth_proxy) -class ProxyConnectTestCase(TestCase): +class TestProxyConnect(TestCase): @classmethod def setUpClass(cls): cls.mockserver = MockServer() @@ -75,7 +76,7 @@ class ProxyConnectTestCase(TestCase): try: import mitmproxy # noqa: F401 except ImportError: - self.skipTest("mitmproxy is not installed") + pytest.skip("mitmproxy is not installed") self._oldenv = os.environ.copy() @@ -113,12 +114,12 @@ class ProxyConnectTestCase(TestCase): yield crawler.crawl(seed=request) self._assert_got_response_code(200, log) echo = json.loads(crawler.spider.meta["responses"][0].text) - self.assertTrue("Proxy-Authorization" not in echo["headers"]) + assert "Proxy-Authorization" not in echo["headers"] def _assert_got_response_code(self, code, log): print(log) - self.assertEqual(str(log).count(f"Crawled ({code})"), 1) + assert str(log).count(f"Crawled ({code})") == 1 def _assert_got_tunnel_error(self, log): print(log) - self.assertIn("TunnelError", str(log)) + assert "TunnelError" in str(log) diff --git a/tests/test_request_attribute_binding.py b/tests/test_request_attribute_binding.py index 0072660a7..9b42fd6c7 100644 --- a/tests/test_request_attribute_binding.py +++ b/tests/test_request_attribute_binding.py @@ -56,7 +56,7 @@ class AlternativeCallbacksMiddleware: return response.replace(request=new_request) -class CrawlTestCase(TestCase): +class TestCrawl(TestCase): @classmethod def setUpClass(cls): cls.mockserver = MockServer() @@ -72,7 +72,7 @@ class CrawlTestCase(TestCase): crawler = get_crawler(SingleRequestSpider) yield crawler.crawl(seed=url, mockserver=self.mockserver) response = crawler.spider.meta["responses"][0] - self.assertEqual(response.request.url, url) + assert response.request.url == url @defer.inlineCallbacks def test_response_error(self): @@ -82,8 +82,8 @@ class CrawlTestCase(TestCase): yield crawler.crawl(seed=url, mockserver=self.mockserver) failure = crawler.spider.meta["failure"] response = failure.value.response - self.assertEqual(failure.request.url, url) - self.assertEqual(response.request.url, url) + assert failure.request.url == url + assert response.request.url == url @defer.inlineCallbacks def test_downloader_middleware_raise_exception(self): @@ -98,8 +98,8 @@ class CrawlTestCase(TestCase): ) yield crawler.crawl(seed=url, mockserver=self.mockserver) failure = crawler.spider.meta["failure"] - self.assertEqual(failure.request.url, url) - self.assertIsInstance(failure.value, ZeroDivisionError) + assert failure.request.url == url + assert isinstance(failure.value, ZeroDivisionError) @defer.inlineCallbacks def test_downloader_middleware_override_request_in_process_response(self): @@ -131,10 +131,10 @@ class CrawlTestCase(TestCase): yield crawler.crawl(seed=url, mockserver=self.mockserver) response = crawler.spider.meta["responses"][0] - self.assertEqual(response.request.url, OVERRIDDEN_URL) + assert response.request.url == OVERRIDDEN_URL - self.assertEqual(signal_params["response"].url, url) - self.assertEqual(signal_params["request"].url, OVERRIDDEN_URL) + assert signal_params["response"].url == url + assert signal_params["request"].url == OVERRIDDEN_URL log.check_present( ( @@ -164,8 +164,8 @@ class CrawlTestCase(TestCase): ) yield crawler.crawl(seed=url, mockserver=self.mockserver) response = crawler.spider.meta["responses"][0] - self.assertEqual(response.body, b"Caught ZeroDivisionError") - self.assertEqual(response.request.url, OVERRIDDEN_URL) + assert response.body == b"Caught ZeroDivisionError" + assert response.request.url == OVERRIDDEN_URL @defer.inlineCallbacks def test_downloader_middleware_do_not_override_in_process_exception(self): @@ -187,8 +187,8 @@ class CrawlTestCase(TestCase): ) yield crawler.crawl(seed=url, mockserver=self.mockserver) response = crawler.spider.meta["responses"][0] - self.assertEqual(response.body, b"Caught ZeroDivisionError") - self.assertEqual(response.request.url, url) + assert response.body == b"Caught ZeroDivisionError" + assert response.request.url == url @defer.inlineCallbacks def test_downloader_middleware_alternative_callback(self): diff --git a/tests/test_request_cb_kwargs.py b/tests/test_request_cb_kwargs.py index a21cb43ff..ab6baa5f0 100644 --- a/tests/test_request_cb_kwargs.py +++ b/tests/test_request_cb_kwargs.py @@ -151,7 +151,7 @@ class KeywordArgumentsSpider(MockServerSpider): self.crawler.stats.inc_value("boolean_checks", 1) -class CallbackKeywordArgumentsTestCase(TestCase): +class TestCallbackKeywordArguments(TestCase): maxDiff = None @classmethod @@ -168,27 +168,19 @@ class CallbackKeywordArgumentsTestCase(TestCase): crawler = get_crawler(KeywordArgumentsSpider) with LogCapture() as log: yield crawler.crawl(mockserver=self.mockserver) - self.assertTrue(all(crawler.spider.checks)) - self.assertEqual( - len(crawler.spider.checks), crawler.stats.get_value("boolean_checks") - ) + assert all(crawler.spider.checks) + assert len(crawler.spider.checks) == crawler.stats.get_value("boolean_checks") # check exceptions for argument mismatch exceptions = {} for line in log.records: for key in ("takes_less", "takes_more"): if key in line.getMessage(): exceptions[key] = line - self.assertEqual(exceptions["takes_less"].exc_info[0], TypeError) - self.assertTrue( - str(exceptions["takes_less"].exc_info[1]).endswith( - "parse_takes_less() got an unexpected keyword argument 'number'" - ), - msg="Exception message: " + str(exceptions["takes_less"].exc_info[1]), - ) - self.assertEqual(exceptions["takes_more"].exc_info[0], TypeError) - self.assertTrue( - str(exceptions["takes_more"].exc_info[1]).endswith( - "parse_takes_more() missing 1 required positional argument: 'other'" - ), - msg="Exception message: " + str(exceptions["takes_more"].exc_info[1]), - ) + assert exceptions["takes_less"].exc_info[0] is TypeError + assert str(exceptions["takes_less"].exc_info[1]).endswith( + "parse_takes_less() got an unexpected keyword argument 'number'" + ), "Exception message: " + str(exceptions["takes_less"].exc_info[1]) + assert exceptions["takes_more"].exc_info[0] is TypeError + assert str(exceptions["takes_more"].exc_info[1]).endswith( + "parse_takes_more() missing 1 required positional argument: 'other'" + ), "Exception message: " + str(exceptions["takes_more"].exc_info[1]) diff --git a/tests/test_request_dict.py b/tests/test_request_dict.py index 2c605a015..ea7018541 100644 --- a/tests/test_request_dict.py +++ b/tests/test_request_dict.py @@ -1,5 +1,3 @@ -import unittest - import pytest from scrapy import Request, Spider @@ -11,8 +9,8 @@ class CustomRequest(Request): pass -class RequestSerializationTest(unittest.TestCase): - def setUp(self): +class TestRequestSerialization: + def setup_method(self): self.spider = MethodsSpider() def test_basic(self): @@ -50,23 +48,23 @@ class RequestSerializationTest(unittest.TestCase): self._assert_same_request(request, request2) def _assert_same_request(self, r1, r2): - self.assertEqual(r1.__class__, r2.__class__) - self.assertEqual(r1.url, r2.url) - self.assertEqual(r1.callback, r2.callback) - self.assertEqual(r1.errback, r2.errback) - self.assertEqual(r1.method, r2.method) - self.assertEqual(r1.body, r2.body) - self.assertEqual(r1.headers, r2.headers) - self.assertEqual(r1.cookies, r2.cookies) - self.assertEqual(r1.meta, r2.meta) - self.assertEqual(r1.cb_kwargs, r2.cb_kwargs) - self.assertEqual(r1.encoding, r2.encoding) - self.assertEqual(r1._encoding, r2._encoding) - self.assertEqual(r1.priority, r2.priority) - self.assertEqual(r1.dont_filter, r2.dont_filter) - self.assertEqual(r1.flags, r2.flags) + assert r1.__class__ == r2.__class__ + assert r1.url == r2.url + assert r1.callback == r2.callback + assert r1.errback == r2.errback + assert r1.method == r2.method + assert r1.body == r2.body + assert r1.headers == r2.headers + assert r1.cookies == r2.cookies + assert r1.meta == r2.meta + assert r1.cb_kwargs == r2.cb_kwargs + assert r1.encoding == r2.encoding + assert r1._encoding == r2._encoding + assert r1.priority == r2.priority + assert r1.dont_filter == r2.dont_filter + assert r1.flags == r2.flags if isinstance(r1, JsonRequest): - self.assertEqual(r1.dumps_kwargs, r2.dumps_kwargs) + assert r1.dumps_kwargs == r2.dumps_kwargs def test_request_class(self): r1 = FormRequest("http://www.example.com") @@ -92,8 +90,8 @@ class RequestSerializationTest(unittest.TestCase): ) self._assert_serializes_ok(r, spider=self.spider) request_dict = r.to_dict(spider=self.spider) - self.assertEqual(request_dict["callback"], "parse_item_reference") - self.assertEqual(request_dict["errback"], "handle_error_reference") + assert request_dict["callback"] == "parse_item_reference" + assert request_dict["errback"] == "handle_error_reference" def test_private_reference_callback_serialization(self): r = Request( @@ -103,12 +101,8 @@ class RequestSerializationTest(unittest.TestCase): ) self._assert_serializes_ok(r, spider=self.spider) request_dict = r.to_dict(spider=self.spider) - self.assertEqual( - request_dict["callback"], "_MethodsSpider__parse_item_reference" - ) - self.assertEqual( - request_dict["errback"], "_MethodsSpider__handle_error_reference" - ) + assert request_dict["callback"] == "_MethodsSpider__parse_item_reference" + assert request_dict["errback"] == "_MethodsSpider__handle_error_reference" def test_private_callback_serialization(self): r = Request( diff --git a/tests/test_request_left.py b/tests/test_request_left.py index cf4c8a2d5..d55905f9c 100644 --- a/tests/test_request_left.py +++ b/tests/test_request_left.py @@ -38,22 +38,22 @@ class TestCatching(TestCase): def test_success(self): crawler = get_crawler(SignalCatcherSpider) yield crawler.crawl(self.mockserver.url("/status?n=200")) - self.assertEqual(crawler.spider.caught_times, 1) + assert crawler.spider.caught_times == 1 @defer.inlineCallbacks def test_timeout(self): crawler = get_crawler(SignalCatcherSpider, {"DOWNLOAD_TIMEOUT": 0.1}) yield crawler.crawl(self.mockserver.url("/delay?n=0.2")) - self.assertEqual(crawler.spider.caught_times, 1) + assert crawler.spider.caught_times == 1 @defer.inlineCallbacks def test_disconnect(self): crawler = get_crawler(SignalCatcherSpider) yield crawler.crawl(self.mockserver.url("/drop")) - self.assertEqual(crawler.spider.caught_times, 1) + assert crawler.spider.caught_times == 1 @defer.inlineCallbacks def test_noconnect(self): crawler = get_crawler(SignalCatcherSpider) yield crawler.crawl("http://thereisdefinetelynosuchdomain.com") - self.assertEqual(crawler.spider.caught_times, 1) + assert crawler.spider.caught_times == 1 diff --git a/tests/test_responsetypes.py b/tests/test_responsetypes.py index f9f56ff97..5b04c7436 100644 --- a/tests/test_responsetypes.py +++ b/tests/test_responsetypes.py @@ -1,5 +1,3 @@ -import unittest - from scrapy.http import ( Headers, HtmlResponse, @@ -11,7 +9,7 @@ from scrapy.http import ( from scrapy.responsetypes import responsetypes -class ResponseTypesTest(unittest.TestCase): +class TestResponseTypes: def test_from_filename(self): mappings = [ ("data.bin", Response), @@ -123,6 +121,4 @@ class ResponseTypesTest(unittest.TestCase): def test_custom_mime_types_loaded(self): # check that mime.types files shipped with scrapy are loaded - self.assertEqual( - responsetypes.mimetypes.guess_type("x.scrapytest")[0], "x-scrapy/test" - ) + assert responsetypes.mimetypes.guess_type("x.scrapytest")[0] == "x-scrapy/test" diff --git a/tests/test_robotstxt_interface.py b/tests/test_robotstxt_interface.py index 0d00ff660..221ccabe6 100644 --- a/tests/test_robotstxt_interface.py +++ b/tests/test_robotstxt_interface.py @@ -1,4 +1,4 @@ -from twisted.trial import unittest +import pytest from scrapy.robotstxt import decode_robotstxt @@ -32,8 +32,8 @@ class BaseRobotParserTest: rp = self.parser_cls.from_crawler( crawler=None, robotstxt_body=robotstxt_robotstxt_body ) - self.assertTrue(rp.allowed("https://www.site.local/allowed", "*")) - self.assertFalse(rp.allowed("https://www.site.local/disallowed", "*")) + assert rp.allowed("https://www.site.local/allowed", "*") + assert not rp.allowed("https://www.site.local/disallowed", "*") def test_allowed_wildcards(self): robotstxt_robotstxt_body = b"""User-agent: first @@ -47,42 +47,36 @@ class BaseRobotParserTest: crawler=None, robotstxt_body=robotstxt_robotstxt_body ) - self.assertTrue(rp.allowed("https://www.site.local/disallowed", "first")) - self.assertFalse( - rp.allowed("https://www.site.local/disallowed/xyz/end", "first") - ) - self.assertFalse( - rp.allowed("https://www.site.local/disallowed/abc/end", "first") - ) - self.assertTrue( - rp.allowed("https://www.site.local/disallowed/xyz/endinglater", "first") - ) + assert rp.allowed("https://www.site.local/disallowed", "first") + assert not rp.allowed("https://www.site.local/disallowed/xyz/end", "first") + assert not rp.allowed("https://www.site.local/disallowed/abc/end", "first") + assert rp.allowed("https://www.site.local/disallowed/xyz/endinglater", "first") - self.assertTrue(rp.allowed("https://www.site.local/allowed", "second")) - self.assertTrue(rp.allowed("https://www.site.local/is_still_allowed", "second")) - self.assertTrue(rp.allowed("https://www.site.local/is_allowed_too", "second")) + assert rp.allowed("https://www.site.local/allowed", "second") + assert rp.allowed("https://www.site.local/is_still_allowed", "second") + assert rp.allowed("https://www.site.local/is_allowed_too", "second") def test_length_based_precedence(self): robotstxt_robotstxt_body = b"User-agent: * \nDisallow: / \nAllow: /page" rp = self.parser_cls.from_crawler( crawler=None, robotstxt_body=robotstxt_robotstxt_body ) - self.assertTrue(rp.allowed("https://www.site.local/page", "*")) + assert rp.allowed("https://www.site.local/page", "*") def test_order_based_precedence(self): robotstxt_robotstxt_body = b"User-agent: * \nDisallow: / \nAllow: /page" rp = self.parser_cls.from_crawler( crawler=None, robotstxt_body=robotstxt_robotstxt_body ) - self.assertFalse(rp.allowed("https://www.site.local/page", "*")) + assert not rp.allowed("https://www.site.local/page", "*") def test_empty_response(self): """empty response should equal 'allow all'""" rp = self.parser_cls.from_crawler(crawler=None, robotstxt_body=b"") - self.assertTrue(rp.allowed("https://site.local/", "*")) - self.assertTrue(rp.allowed("https://site.local/", "chrome")) - self.assertTrue(rp.allowed("https://site.local/index.html", "*")) - self.assertTrue(rp.allowed("https://site.local/disallowed", "*")) + assert rp.allowed("https://site.local/", "*") + assert rp.allowed("https://site.local/", "chrome") + assert rp.allowed("https://site.local/index.html", "*") + assert rp.allowed("https://site.local/disallowed", "*") def test_garbage_response(self): """garbage response should be discarded, equal 'allow all'""" @@ -90,10 +84,10 @@ class BaseRobotParserTest: rp = self.parser_cls.from_crawler( crawler=None, robotstxt_body=robotstxt_robotstxt_body ) - self.assertTrue(rp.allowed("https://site.local/", "*")) - self.assertTrue(rp.allowed("https://site.local/", "chrome")) - self.assertTrue(rp.allowed("https://site.local/index.html", "*")) - self.assertTrue(rp.allowed("https://site.local/disallowed", "*")) + assert rp.allowed("https://site.local/", "*") + assert rp.allowed("https://site.local/", "chrome") + assert rp.allowed("https://site.local/index.html", "*") + assert rp.allowed("https://site.local/disallowed", "*") def test_unicode_url_and_useragent(self): robotstxt_robotstxt_body = """ @@ -109,79 +103,67 @@ class BaseRobotParserTest: rp = self.parser_cls.from_crawler( crawler=None, robotstxt_body=robotstxt_robotstxt_body ) - self.assertTrue(rp.allowed("https://site.local/", "*")) - self.assertFalse(rp.allowed("https://site.local/admin/", "*")) - self.assertFalse(rp.allowed("https://site.local/static/", "*")) - self.assertTrue(rp.allowed("https://site.local/admin/", "UnicödeBöt")) - self.assertFalse( - rp.allowed("https://site.local/wiki/K%C3%A4ytt%C3%A4j%C3%A4:", "*") - ) - self.assertFalse(rp.allowed("https://site.local/wiki/Käyttäjä:", "*")) - self.assertTrue(rp.allowed("https://site.local/some/randome/page.html", "*")) - self.assertFalse( - rp.allowed("https://site.local/some/randome/page.html", "UnicödeBöt") - ) + assert rp.allowed("https://site.local/", "*") + assert not rp.allowed("https://site.local/admin/", "*") + assert not rp.allowed("https://site.local/static/", "*") + assert rp.allowed("https://site.local/admin/", "UnicödeBöt") + assert not rp.allowed("https://site.local/wiki/K%C3%A4ytt%C3%A4j%C3%A4:", "*") + assert not rp.allowed("https://site.local/wiki/Käyttäjä:", "*") + assert rp.allowed("https://site.local/some/randome/page.html", "*") + assert not rp.allowed("https://site.local/some/randome/page.html", "UnicödeBöt") -class DecodeRobotsTxtTest(unittest.TestCase): +class TestDecodeRobotsTxt: def test_native_string_conversion(self): robotstxt_body = b"User-agent: *\nDisallow: /\n" decoded_content = decode_robotstxt( robotstxt_body, spider=None, to_native_str_type=True ) - self.assertEqual(decoded_content, "User-agent: *\nDisallow: /\n") + assert decoded_content == "User-agent: *\nDisallow: /\n" def test_decode_utf8(self): robotstxt_body = b"User-agent: *\nDisallow: /\n" decoded_content = decode_robotstxt(robotstxt_body, spider=None) - self.assertEqual(decoded_content, "User-agent: *\nDisallow: /\n") + assert decoded_content == "User-agent: *\nDisallow: /\n" def test_decode_non_utf8(self): robotstxt_body = b"User-agent: *\n\xffDisallow: /\n" decoded_content = decode_robotstxt(robotstxt_body, spider=None) - self.assertEqual(decoded_content, "User-agent: *\nDisallow: /\n") + assert decoded_content == "User-agent: *\nDisallow: /\n" -class PythonRobotParserTest(BaseRobotParserTest, unittest.TestCase): - def setUp(self): +class TestPythonRobotParser(BaseRobotParserTest): + def setup_method(self): from scrapy.robotstxt import PythonRobotParser super()._setUp(PythonRobotParser) def test_length_based_precedence(self): - raise unittest.SkipTest( + pytest.skip( "RobotFileParser does not support length based directives precedence." ) def test_allowed_wildcards(self): - raise unittest.SkipTest("RobotFileParser does not support wildcards.") + pytest.skip("RobotFileParser does not support wildcards.") -class RerpRobotParserTest(BaseRobotParserTest, unittest.TestCase): - if not rerp_available(): - skip = "Rerp parser is not installed" - - def setUp(self): +@pytest.mark.skipif(not rerp_available(), reason="Rerp parser is not installed") +class TestRerpRobotParser(BaseRobotParserTest): + def setup_method(self): from scrapy.robotstxt import RerpRobotParser super()._setUp(RerpRobotParser) def test_length_based_precedence(self): - raise unittest.SkipTest( - "Rerp does not support length based directives precedence." - ) + pytest.skip("Rerp does not support length based directives precedence.") -class ProtegoRobotParserTest(BaseRobotParserTest, unittest.TestCase): - if not protego_available(): - skip = "Protego parser is not installed" - - def setUp(self): +@pytest.mark.skipif(not protego_available(), reason="Protego parser is not installed") +class TestProtegoRobotParser(BaseRobotParserTest): + def setup_method(self): from scrapy.robotstxt import ProtegoRobotParser super()._setUp(ProtegoRobotParser) def test_order_based_precedence(self): - raise unittest.SkipTest( - "Protego does not support order based directives precedence." - ) + pytest.skip("Protego does not support order based directives precedence.") diff --git a/tests/test_scheduler.py b/tests/test_scheduler.py index f2f8b96cd..1d6992a32 100644 --- a/tests/test_scheduler.py +++ b/tests/test_scheduler.py @@ -2,7 +2,7 @@ from __future__ import annotations import shutil import tempfile -import unittest +from abc import ABC, abstractmethod from typing import Any, NamedTuple import pytest @@ -65,10 +65,14 @@ class MockCrawler(Crawler): self.stats = load_object(self.settings["STATS_CLASS"])(self) -class SchedulerHandler: - priority_queue_cls: str | None = None +class SchedulerHandler(ABC): jobdir = None + @property + @abstractmethod + def priority_queue_cls(self) -> str: + raise NotImplementedError + def create_scheduler(self): self.mock_crawler = MockCrawler(self.priority_queue_cls, self.jobdir) self.scheduler = Scheduler.from_crawler(self.mock_crawler) @@ -80,10 +84,10 @@ class SchedulerHandler: self.mock_crawler.stop() self.mock_crawler.engine.downloader.close() - def setUp(self): + def setup_method(self): self.create_scheduler() - def tearDown(self): + def teardown_method(self): self.close_scheduler() @@ -99,16 +103,16 @@ _PRIORITIES = [ _URLS = {"http://foo.com/a", "http://foo.com/b", "http://foo.com/c"} -class BaseSchedulerInMemoryTester(SchedulerHandler): +class TestSchedulerInMemoryBase(SchedulerHandler): def test_length(self): - self.assertFalse(self.scheduler.has_pending_requests()) - self.assertEqual(len(self.scheduler), 0) + assert not self.scheduler.has_pending_requests() + assert len(self.scheduler) == 0 for url in _URLS: self.scheduler.enqueue_request(Request(url)) - self.assertTrue(self.scheduler.has_pending_requests()) - self.assertEqual(len(self.scheduler), len(_URLS)) + assert self.scheduler.has_pending_requests() + assert len(self.scheduler) == len(_URLS) def test_dequeue(self): for url in _URLS: @@ -118,7 +122,7 @@ class BaseSchedulerInMemoryTester(SchedulerHandler): while self.scheduler.has_pending_requests(): urls.add(self.scheduler.next_request().url) - self.assertEqual(urls, _URLS) + assert urls == _URLS def test_dequeue_priorities(self): for url, priority in _PRIORITIES: @@ -128,25 +132,23 @@ class BaseSchedulerInMemoryTester(SchedulerHandler): while self.scheduler.has_pending_requests(): priorities.append(self.scheduler.next_request().priority) - self.assertEqual( - priorities, sorted([x[1] for x in _PRIORITIES], key=lambda x: -x) - ) + assert priorities == sorted([x[1] for x in _PRIORITIES], key=lambda x: -x) -class BaseSchedulerOnDiskTester(SchedulerHandler): - def setUp(self): +class TestSchedulerOnDiskBase(SchedulerHandler): + def setup_method(self): self.jobdir = tempfile.mkdtemp() self.create_scheduler() - def tearDown(self): + def teardown_method(self): self.close_scheduler() shutil.rmtree(self.jobdir) self.jobdir = None def test_length(self): - self.assertFalse(self.scheduler.has_pending_requests()) - self.assertEqual(len(self.scheduler), 0) + assert not self.scheduler.has_pending_requests() + assert len(self.scheduler) == 0 for url in _URLS: self.scheduler.enqueue_request(Request(url)) @@ -154,8 +156,8 @@ class BaseSchedulerOnDiskTester(SchedulerHandler): self.close_scheduler() self.create_scheduler() - self.assertTrue(self.scheduler.has_pending_requests()) - self.assertEqual(len(self.scheduler), len(_URLS)) + assert self.scheduler.has_pending_requests() + assert len(self.scheduler) == len(_URLS) def test_dequeue(self): for url in _URLS: @@ -168,7 +170,7 @@ class BaseSchedulerOnDiskTester(SchedulerHandler): while self.scheduler.has_pending_requests(): urls.add(self.scheduler.next_request().url) - self.assertEqual(urls, _URLS) + assert urls == _URLS def test_dequeue_priorities(self): for url, priority in _PRIORITIES: @@ -181,17 +183,19 @@ class BaseSchedulerOnDiskTester(SchedulerHandler): while self.scheduler.has_pending_requests(): priorities.append(self.scheduler.next_request().priority) - self.assertEqual( - priorities, sorted([x[1] for x in _PRIORITIES], key=lambda x: -x) - ) + assert priorities == sorted([x[1] for x in _PRIORITIES], key=lambda x: -x) -class TestSchedulerInMemory(BaseSchedulerInMemoryTester, unittest.TestCase): - priority_queue_cls = "scrapy.pqueues.ScrapyPriorityQueue" +class TestSchedulerInMemory(TestSchedulerInMemoryBase): + @property + def priority_queue_cls(self) -> str: + return "scrapy.pqueues.ScrapyPriorityQueue" -class TestSchedulerOnDisk(BaseSchedulerOnDiskTester, unittest.TestCase): - priority_queue_cls = "scrapy.pqueues.ScrapyPriorityQueue" +class TestSchedulerOnDisk(TestSchedulerOnDiskBase): + @property + def priority_queue_cls(self) -> str: + return "scrapy.pqueues.ScrapyPriorityQueue" _URLS_WITH_SLOTS = [ @@ -204,37 +208,34 @@ _URLS_WITH_SLOTS = [ ] -class TestMigration(unittest.TestCase): - def setUp(self): - self.tmpdir = tempfile.mkdtemp() +class TestMigration: + def test_migration(self, tmpdir): + class PrevSchedulerHandler(SchedulerHandler): + jobdir = tmpdir - def tearDown(self): - shutil.rmtree(self.tmpdir) + @property + def priority_queue_cls(self) -> str: + return "scrapy.pqueues.ScrapyPriorityQueue" - def _migration(self, tmp_dir): - prev_scheduler_handler = SchedulerHandler() - prev_scheduler_handler.priority_queue_cls = "scrapy.pqueues.ScrapyPriorityQueue" - prev_scheduler_handler.jobdir = tmp_dir + class NextSchedulerHandler(SchedulerHandler): + jobdir = tmpdir + @property + def priority_queue_cls(self) -> str: + return "scrapy.pqueues.DownloaderAwarePriorityQueue" + + prev_scheduler_handler = PrevSchedulerHandler() prev_scheduler_handler.create_scheduler() for url in _URLS: prev_scheduler_handler.scheduler.enqueue_request(Request(url)) prev_scheduler_handler.close_scheduler() - next_scheduler_handler = SchedulerHandler() - next_scheduler_handler.priority_queue_cls = ( - "scrapy.pqueues.DownloaderAwarePriorityQueue" - ) - next_scheduler_handler.jobdir = tmp_dir - - next_scheduler_handler.create_scheduler() - - def test_migration(self): + next_scheduler_handler = NextSchedulerHandler() with pytest.raises( ValueError, match="DownloaderAwarePriorityQueue accepts ``slot_startprios`` as a dict", ): - self._migration(self.tmpdir) + next_scheduler_handler.create_scheduler() def _is_scheduling_fair(enqueued_slots, dequeued_slots): @@ -263,9 +264,12 @@ def _is_scheduling_fair(enqueued_slots, dequeued_slots): class DownloaderAwareSchedulerTestMixin: - priority_queue_cls: str | None = "scrapy.pqueues.DownloaderAwarePriorityQueue" reopen = False + @property + def priority_queue_cls(self) -> str: + return "scrapy.pqueues.DownloaderAwarePriorityQueue" + def test_logic(self): for url, slot in _URLS_WITH_SLOTS: request = Request(url) @@ -290,20 +294,18 @@ class DownloaderAwareSchedulerTestMixin: slot = downloader.get_slot_key(request) downloader.decrement(slot) - self.assertTrue( - _is_scheduling_fair([s for u, s in _URLS_WITH_SLOTS], dequeued_slots) - ) - self.assertEqual(sum(len(s.active) for s in downloader.slots.values()), 0) + assert _is_scheduling_fair([s for u, s in _URLS_WITH_SLOTS], dequeued_slots) + assert sum(len(s.active) for s in downloader.slots.values()) == 0 class TestSchedulerWithDownloaderAwareInMemory( - DownloaderAwareSchedulerTestMixin, BaseSchedulerInMemoryTester, unittest.TestCase + DownloaderAwareSchedulerTestMixin, TestSchedulerInMemoryBase ): pass class TestSchedulerWithDownloaderAwareOnDisk( - DownloaderAwareSchedulerTestMixin, BaseSchedulerOnDiskTester, unittest.TestCase + DownloaderAwareSchedulerTestMixin, TestSchedulerOnDiskBase ): reopen = True @@ -337,13 +339,12 @@ class TestIntegrationWithDownloaderAwareInMemory(TestCase): url = mockserver.url("/status?n=200", is_secure=False) start_urls = [url] * 6 yield self.crawler.crawl(start_urls) - self.assertEqual( - self.crawler.stats.get_value("downloader/response_count"), - len(start_urls), + assert self.crawler.stats.get_value("downloader/response_count") == len( + start_urls ) -class TestIncompatibility(unittest.TestCase): +class TestIncompatibility: def _incompatible(self): settings = { "SCHEDULER_PRIORITY_QUEUE": "scrapy.pqueues.DownloaderAwarePriorityQueue", diff --git a/tests/test_scheduler_base.py b/tests/test_scheduler_base.py index c2bb8cec5..4a36d3cdb 100644 --- a/tests/test_scheduler_base.py +++ b/tests/test_scheduler_base.py @@ -1,12 +1,11 @@ from __future__ import annotations -from unittest import TestCase from urllib.parse import urljoin import pytest from testfixtures import LogCapture from twisted.internet import defer -from twisted.trial.unittest import TestCase as TwistedTestCase +from twisted.trial.unittest import TestCase from scrapy.core.scheduler import BaseScheduler from scrapy.http import Request @@ -65,17 +64,17 @@ class PathsSpider(Spider): class InterfaceCheckMixin: def test_scheduler_class(self): - self.assertTrue(isinstance(self.scheduler, BaseScheduler)) - self.assertTrue(issubclass(self.scheduler.__class__, BaseScheduler)) + assert isinstance(self.scheduler, BaseScheduler) + assert issubclass(self.scheduler.__class__, BaseScheduler) -class BaseSchedulerTest(TestCase, InterfaceCheckMixin): - def setUp(self): +class TestBaseScheduler(InterfaceCheckMixin): + def setup_method(self): self.scheduler = BaseScheduler() def test_methods(self): - self.assertIsNone(self.scheduler.open(Spider("foo"))) - self.assertIsNone(self.scheduler.close("finished")) + assert self.scheduler.open(Spider("foo")) is None + assert self.scheduler.close("finished") is None with pytest.raises(NotImplementedError): self.scheduler.has_pending_requests() with pytest.raises(NotImplementedError): @@ -84,8 +83,8 @@ class BaseSchedulerTest(TestCase, InterfaceCheckMixin): self.scheduler.next_request() -class MinimalSchedulerTest(TestCase, InterfaceCheckMixin): - def setUp(self): +class TestMinimalScheduler(InterfaceCheckMixin): + def setup_method(self): self.scheduler = MinimalScheduler() def test_open_close(self): @@ -101,51 +100,51 @@ class MinimalSchedulerTest(TestCase, InterfaceCheckMixin): len(self.scheduler) def test_enqueue_dequeue(self): - self.assertFalse(self.scheduler.has_pending_requests()) + assert not self.scheduler.has_pending_requests() for url in URLS: - self.assertTrue(self.scheduler.enqueue_request(Request(url))) - self.assertFalse(self.scheduler.enqueue_request(Request(url))) - self.assertTrue(self.scheduler.has_pending_requests) + assert self.scheduler.enqueue_request(Request(url)) + assert not self.scheduler.enqueue_request(Request(url)) + assert self.scheduler.has_pending_requests dequeued = [] while self.scheduler.has_pending_requests(): request = self.scheduler.next_request() dequeued.append(request.url) - self.assertEqual(set(dequeued), set(URLS)) - self.assertFalse(self.scheduler.has_pending_requests()) + assert set(dequeued) == set(URLS) + assert not self.scheduler.has_pending_requests() -class SimpleSchedulerTest(TwistedTestCase, InterfaceCheckMixin): +class SimpleSchedulerTest(TestCase, InterfaceCheckMixin): def setUp(self): self.scheduler = SimpleScheduler() @defer.inlineCallbacks def test_enqueue_dequeue(self): open_result = yield self.scheduler.open(Spider("foo")) - self.assertEqual(open_result, "open") - self.assertFalse(self.scheduler.has_pending_requests()) + assert open_result == "open" + assert not self.scheduler.has_pending_requests() for url in URLS: - self.assertTrue(self.scheduler.enqueue_request(Request(url))) - self.assertFalse(self.scheduler.enqueue_request(Request(url))) + assert self.scheduler.enqueue_request(Request(url)) + assert not self.scheduler.enqueue_request(Request(url)) - self.assertTrue(self.scheduler.has_pending_requests()) - self.assertEqual(len(self.scheduler), len(URLS)) + assert self.scheduler.has_pending_requests() + assert len(self.scheduler) == len(URLS) dequeued = [] while self.scheduler.has_pending_requests(): request = self.scheduler.next_request() dequeued.append(request.url) - self.assertEqual(set(dequeued), set(URLS)) + assert set(dequeued) == set(URLS) - self.assertFalse(self.scheduler.has_pending_requests()) - self.assertEqual(len(self.scheduler), 0) + assert not self.scheduler.has_pending_requests() + assert len(self.scheduler) == 0 close_result = yield self.scheduler.close("") - self.assertEqual(close_result, "close") + assert close_result == "close" -class MinimalSchedulerCrawlTest(TwistedTestCase): +class MinimalSchedulerCrawlTest(TestCase): scheduler_cls = MinimalScheduler @defer.inlineCallbacks @@ -158,8 +157,8 @@ class MinimalSchedulerCrawlTest(TwistedTestCase): crawler = get_crawler(PathsSpider, settings) yield crawler.crawl(mockserver) for path in PATHS: - self.assertIn(f"{{'path': '{path}'}}", str(log)) - self.assertIn(f"'item_scraped_count': {len(PATHS)}", str(log)) + assert f"{{'path': '{path}'}}" in str(log) + assert f"'item_scraped_count': {len(PATHS)}" in str(log) class SimpleSchedulerCrawlTest(MinimalSchedulerCrawlTest): diff --git a/tests/test_selector.py b/tests/test_selector.py index 2d7a1442e..5c8eadf0b 100644 --- a/tests/test_selector.py +++ b/tests/test_selector.py @@ -3,7 +3,6 @@ import weakref import parsel import pytest from packaging import version -from twisted.trial import unittest from scrapy.http import HtmlResponse, TextResponse, XmlResponse from scrapy.selector import Selector @@ -12,7 +11,7 @@ PARSEL_VERSION = version.parse(getattr(parsel, "__version__", "0.0")) PARSEL_18_PLUS = PARSEL_VERSION >= version.parse("1.8.0") -class SelectorTestCase(unittest.TestCase): +class TestSelector: def test_simple_selection(self): """Simple selector tests""" body = b"

" @@ -20,57 +19,46 @@ class SelectorTestCase(unittest.TestCase): sel = Selector(response) xl = sel.xpath("//input") - self.assertEqual(2, len(xl)) + assert len(xl) == 2 for x in xl: assert isinstance(x, Selector) - self.assertEqual( - sel.xpath("//input").getall(), [x.get() for x in sel.xpath("//input")] - ) - self.assertEqual( - [x.get() for x in sel.xpath("//input[@name='a']/@name")], ["a"] - ) - self.assertEqual( - [ - x.get() - for x in sel.xpath( - "number(concat(//input[@name='a']/@value, //input[@name='b']/@value))" - ) - ], - ["12.0"], - ) - self.assertEqual(sel.xpath("concat('xpath', 'rules')").getall(), ["xpathrules"]) - self.assertEqual( - [ - x.get() - for x in sel.xpath( - "concat(//input[@name='a']/@value, //input[@name='b']/@value)" - ) - ], - ["12"], - ) + assert sel.xpath("//input").getall() == [x.get() for x in sel.xpath("//input")] + assert [x.get() for x in sel.xpath("//input[@name='a']/@name")] == ["a"] + assert [ + x.get() + for x in sel.xpath( + "number(concat(//input[@name='a']/@value, //input[@name='b']/@value))" + ) + ] == ["12.0"] + assert sel.xpath("concat('xpath', 'rules')").getall() == ["xpathrules"] + assert [ + x.get() + for x in sel.xpath( + "concat(//input[@name='a']/@value, //input[@name='b']/@value)" + ) + ] == ["12"] def test_root_base_url(self): body = b'
' url = "http://example.com" response = TextResponse(url=url, body=body, encoding="utf-8") sel = Selector(response) - self.assertEqual(url, sel.root.base) + assert url == sel.root.base def test_flavor_detection(self): text = b'

Hello

' sel = Selector(XmlResponse("http://example.com", body=text, encoding="utf-8")) - self.assertEqual(sel.type, "xml") - self.assertEqual( - sel.xpath("//div").getall(), - ['

Hello

'], - ) + assert sel.type == "xml" + assert sel.xpath("//div").getall() == [ + '

Hello

' + ] sel = Selector(HtmlResponse("http://example.com", body=text, encoding="utf-8")) - self.assertEqual(sel.type, "html") - self.assertEqual( - sel.xpath("//div").getall(), ['

Hello

'] - ) + assert sel.type == "html" + assert sel.xpath("//div").getall() == [ + '

Hello

' + ] def test_http_header_encoding_precedence(self): # '\xa3' = pound symbol in unicode @@ -92,7 +80,7 @@ class SelectorTestCase(unittest.TestCase): url="http://example.com", headers=headers, body=html_utf8 ) x = Selector(response) - self.assertEqual(x.xpath("//span[@id='blank']/text()").getall(), ["\xa3"]) + assert x.xpath("//span[@id='blank']/text()").getall() == ["\xa3"] def test_badly_encoded_body(self): # \xe9 alone isn't valid utf8 sequence @@ -116,7 +104,7 @@ class SelectorTestCase(unittest.TestCase): Selector(TextResponse(url="http://example.com", body=b""), text="") -class JMESPathTestCase(unittest.TestCase): +class TestJMESPath: @pytest.mark.skipif( not PARSEL_18_PLUS, reason="parsel < 1.8 doesn't support jmespath" ) @@ -149,16 +137,13 @@ class JMESPathTestCase(unittest.TestCase): } """ resp = TextResponse(url="http://example.com", body=body, encoding="utf-8") - self.assertEqual( - resp.jmespath("html").get(), - "
a
b
c
def
", + assert ( + resp.jmespath("html").get() + == "
a
b
c
def
" ) - self.assertEqual( - resp.jmespath("html").xpath("//div/a/text()").getall(), - ["a", "b", "d"], - ) - self.assertEqual(resp.jmespath("html").css("div > b").getall(), ["f"]) - self.assertEqual(resp.jmespath("content").jmespath("name.age").get(), "18") + assert resp.jmespath("html").xpath("//div/a/text()").getall() == ["a", "b", "d"] + assert resp.jmespath("html").css("div > b").getall() == ["f"] + assert resp.jmespath("content").jmespath("name.age").get() == "18" @pytest.mark.skipif( not PARSEL_18_PLUS, reason="parsel < 1.8 doesn't support jmespath" @@ -194,15 +179,19 @@ class JMESPathTestCase(unittest.TestCase):
""" resp = TextResponse(url="http://example.com", body=body, encoding="utf-8") - self.assertEqual( - resp.xpath("//div/content/text()").jmespath("user[*].name").getall(), - ["A", "B", "C", "D"], - ) - self.assertEqual( - resp.xpath("//div/content").jmespath("user[*].name").getall(), - ["A", "B", "C", "D"], - ) - self.assertEqual(resp.xpath("//div/content").jmespath("total").get(), "4") + assert resp.xpath("//div/content/text()").jmespath("user[*].name").getall() == [ + "A", + "B", + "C", + "D", + ] + assert resp.xpath("//div/content").jmespath("user[*].name").getall() == [ + "A", + "B", + "C", + "D", + ] + assert resp.xpath("//div/content").jmespath("total").get() == "4" @pytest.mark.skipif( not PARSEL_18_PLUS, reason="parsel < 1.8 doesn't support jmespath" @@ -238,30 +227,26 @@ class JMESPathTestCase(unittest.TestCase):
""" resp = TextResponse(url="http://example.com", body=body, encoding="utf-8") - self.assertEqual( - resp.xpath("//div/content/text()").jmespath("user[*].name").re(r"(\w+)"), - ["A", "B", "C", "D"], - ) - self.assertEqual( - resp.xpath("//div/content").jmespath("user[*].name").re(r"(\w+)"), - ["A", "B", "C", "D"], + assert resp.xpath("//div/content/text()").jmespath("user[*].name").re( + r"(\w+)" + ) == ["A", "B", "C", "D"] + assert resp.xpath("//div/content").jmespath("user[*].name").re(r"(\w+)") == [ + "A", + "B", + "C", + "D", + ] + + assert resp.xpath("//div/content").jmespath("unavailable").re(r"(\d+)") == [] + + assert ( + resp.xpath("//div/content").jmespath("unavailable").re_first(r"(\d+)") + is None ) - self.assertEqual( - resp.xpath("//div/content").jmespath("unavailable").re(r"(\d+)"), [] - ) - - self.assertEqual( - resp.xpath("//div/content").jmespath("unavailable").re_first(r"(\d+)"), - None, - ) - - self.assertEqual( - resp.xpath("//div/content") - .jmespath("user[*].age.to_string(@)") - .re(r"(\d+)"), - ["18", "32", "22", "25"], - ) + assert resp.xpath("//div/content").jmespath("user[*].age.to_string(@)").re( + r"(\d+)" + ) == ["18", "32", "22", "25"] @pytest.mark.skipif(PARSEL_18_PLUS, reason="parsel >= 1.8 supports jmespath") def test_jmespath_not_available(self) -> None: diff --git a/tests/test_signals.py b/tests/test_signals.py index a508eb41a..f5075fb60 100644 --- a/tests/test_signals.py +++ b/tests/test_signals.py @@ -20,7 +20,7 @@ class ItemSpider(Spider): return {"index": response.meta["index"]} -class AsyncSignalTestCase(unittest.TestCase): +class TestAsyncSignal(unittest.TestCase): @classmethod def setUpClass(cls): cls.mockserver = MockServer() @@ -43,6 +43,6 @@ class AsyncSignalTestCase(unittest.TestCase): crawler = get_crawler(ItemSpider) crawler.signals.connect(self._on_item_scraped, signals.item_scraped) yield crawler.crawl(mockserver=self.mockserver) - self.assertEqual(len(self.items), 10) + assert len(self.items) == 10 for index in range(10): - self.assertIn({"index": index}, self.items) + assert {"index": index} in self.items diff --git a/tests/test_toplevel.py b/tests/test_toplevel.py index d272101b8..a4f31096e 100644 --- a/tests/test_toplevel.py +++ b/tests/test_toplevel.py @@ -1,33 +1,31 @@ -from unittest import TestCase - import scrapy -class ToplevelTestCase(TestCase): +class TestToplevel: def test_version(self): - self.assertIs(type(scrapy.__version__), str) + assert isinstance(scrapy.__version__, str) def test_version_info(self): - self.assertIs(type(scrapy.version_info), tuple) + assert isinstance(scrapy.version_info, tuple) def test_request_shortcut(self): from scrapy.http import FormRequest, Request - self.assertIs(scrapy.Request, Request) - self.assertIs(scrapy.FormRequest, FormRequest) + assert scrapy.Request is Request + assert scrapy.FormRequest is FormRequest def test_spider_shortcut(self): from scrapy.spiders import Spider - self.assertIs(scrapy.Spider, Spider) + assert scrapy.Spider is Spider def test_selector_shortcut(self): from scrapy.selector import Selector - self.assertIs(scrapy.Selector, Selector) + assert scrapy.Selector is Selector def test_item_shortcut(self): from scrapy.item import Field, Item - self.assertIs(scrapy.Item, Item) - self.assertIs(scrapy.Field, Field) + assert scrapy.Item is Item + assert scrapy.Field is Field diff --git a/tests/test_urlparse_monkeypatches.py b/tests/test_urlparse_monkeypatches.py index c695968d7..0e1e89e81 100644 --- a/tests/test_urlparse_monkeypatches.py +++ b/tests/test_urlparse_monkeypatches.py @@ -1,11 +1,10 @@ -import unittest from urllib.parse import urlparse -class UrlparseTestCase(unittest.TestCase): +class TestUrlparse: def test_s3_url(self): p = urlparse("s3://bucket/key/name?param=value") - self.assertEqual(p.scheme, "s3") - self.assertEqual(p.hostname, "bucket") - self.assertEqual(p.path, "/key/name") - self.assertEqual(p.query, "param=value") + assert p.scheme == "s3" + assert p.hostname == "bucket" + assert p.path == "/key/name" + assert p.query == "param=value"