Fix overridable methods in MediaPipeline (#6368)

This commit is contained in:
Sanchay Kumar 2024-05-28 14:12:58 +05:30 committed by GitHub
parent 986d1ee1dd
commit cadb0dd707
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 73 additions and 126 deletions

View File

@ -2,6 +2,7 @@ from __future__ import annotations
import functools
import logging
from abc import ABC, abstractmethod
from collections import defaultdict
from typing import TYPE_CHECKING
@ -27,7 +28,7 @@ def _DUMMY_CALLBACK(response):
return response
class MediaPipeline:
class MediaPipeline(ABC):
LOG_FAILED_RESULTS = True
class SpiderInfo:
@ -55,14 +56,6 @@ class MediaPipeline:
self.handle_httpstatus_list = SequenceExclude(range(300, 400))
def _key_for_pipe(self, key, base_class_name=None, settings=None):
"""
>>> MediaPipeline()._key_for_pipe("IMAGES")
'IMAGES'
>>> class MyPipe(MediaPipeline):
... pass
>>> MyPipe()._key_for_pipe("IMAGES", base_class_name="MediaPipeline")
'MYPIPE_IMAGES'
"""
class_name = self.__class__.__name__
formatted_key = f"{class_name.upper()}_{key}"
if (
@ -192,21 +185,25 @@ class MediaPipeline:
defer_result(result).chainDeferred(wad)
# Overridable Interface
@abstractmethod
def media_to_download(self, request, info, *, item=None):
"""Check request before starting download"""
pass
raise NotImplementedError()
@abstractmethod
def get_media_requests(self, item, info):
"""Returns the media requests to download"""
pass
raise NotImplementedError()
@abstractmethod
def media_downloaded(self, response, request, info, *, item=None):
"""Handler for success downloads"""
return response
raise NotImplementedError()
@abstractmethod
def media_failed(self, failure, request, info):
"""Handler for failed downloads"""
return failure
raise NotImplementedError()
def item_completed(self, results, item, info):
"""Called per item when all media requests has been processed"""
@ -221,6 +218,7 @@ class MediaPipeline:
)
return item
@abstractmethod
def file_path(self, request, response=None, info=None, *, item=None):
"""Returns the path where downloaded media should be stored"""
pass
raise NotImplementedError()

View File

@ -1,4 +1,3 @@
import io
from typing import Optional
from testfixtures import LogCapture
@ -11,7 +10,6 @@ from scrapy import signals
from scrapy.http import Request, Response
from scrapy.http.request import NO_CALLBACK
from scrapy.pipelines.files import FileException
from scrapy.pipelines.images import ImagesPipeline
from scrapy.pipelines.media import MediaPipeline
from scrapy.settings import Settings
from scrapy.spiders import Spider
@ -35,8 +33,26 @@ def _mocked_download_func(request, info):
return response() if callable(response) else response
class UserDefinedPipeline(MediaPipeline):
def media_to_download(self, request, info, *, item=None):
pass
def get_media_requests(self, item, info):
pass
def media_downloaded(self, response, request, info, *, item=None):
return {}
def media_failed(self, failure, request, info):
return failure
def file_path(self, request, response=None, info=None, *, item=None):
return ""
class BaseMediaPipelineTestCase(unittest.TestCase):
pipeline_class = MediaPipeline
pipeline_class = UserDefinedPipeline
settings = None
def setUp(self):
@ -54,54 +70,6 @@ class BaseMediaPipelineTestCase(unittest.TestCase):
if not name.startswith("_"):
disconnect_all(signal)
def test_default_media_to_download(self):
request = Request("http://url")
assert self.pipe.media_to_download(request, self.info) is None
def test_default_get_media_requests(self):
item = {"name": "name"}
assert self.pipe.get_media_requests(item, self.info) is None
def test_default_media_downloaded(self):
request = Request("http://url")
response = Response("http://url", body=b"")
assert self.pipe.media_downloaded(response, request, self.info) is response
def test_default_media_failed(self):
request = Request("http://url")
fail = Failure(Exception())
assert self.pipe.media_failed(fail, request, self.info) is fail
def test_default_item_completed(self):
item = {"name": "name"}
assert self.pipe.item_completed([], item, self.info) is item
# Check that failures are logged by default
fail = Failure(Exception())
results = [(True, 1), (False, fail)]
with LogCapture() as log:
new_item = self.pipe.item_completed(results, item, self.info)
assert new_item is item
assert len(log.records) == 1
record = log.records[0]
assert record.levelname == "ERROR"
self.assertTupleEqual(record.exc_info, failure_to_exc_info(fail))
# disable failure logging and check again
self.pipe.LOG_FAILED_RESULTS = False
with LogCapture() as log:
new_item = self.pipe.item_completed(results, item, self.info)
assert new_item is item
assert len(log.records) == 0
@inlineCallbacks
def test_default_process_item(self):
item = {"name": "name"}
new_item = yield self.pipe.process_item(item, self.spider)
assert new_item is item
def test_modify_media_request(self):
request = Request("http://url")
self.pipe._modify_media_request(request)
@ -175,8 +143,38 @@ class BaseMediaPipelineTestCase(unittest.TestCase):
context = getattr(info.downloaded[fp].value, "__context__", None)
self.assertIsNone(context)
def test_default_item_completed(self):
item = {"name": "name"}
assert self.pipe.item_completed([], item, self.info) is item
class MockedMediaPipeline(MediaPipeline):
# Check that failures are logged by default
fail = Failure(Exception())
results = [(True, 1), (False, fail)]
with LogCapture() as log:
new_item = self.pipe.item_completed(results, item, self.info)
assert new_item is item
assert len(log.records) == 1
record = log.records[0]
assert record.levelname == "ERROR"
self.assertTupleEqual(record.exc_info, failure_to_exc_info(fail))
# disable failure logging and check again
self.pipe.LOG_FAILED_RESULTS = False
with LogCapture() as log:
new_item = self.pipe.item_completed(results, item, self.info)
assert new_item is item
assert len(log.records) == 0
@inlineCallbacks
def test_default_process_item(self):
item = {"name": "name"}
new_item = yield self.pipe.process_item(item, self.spider)
assert new_item is item
class MockedMediaPipeline(UserDefinedPipeline):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self._mockcalled = []
@ -232,7 +230,7 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase):
)
item = {"requests": req}
new_item = yield self.pipe.process_item(item, self.spider)
self.assertEqual(new_item["results"], [(True, rsp)])
self.assertEqual(new_item["results"], [(True, {})])
self.assertEqual(
self.pipe._mockcalled,
[
@ -277,7 +275,7 @@ 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, rsp1), (False, fail)])
self.assertEqual(new_item["results"], [(True, {}), (False, fail)])
m = self.pipe._mockcalled
# only once
self.assertEqual(m[0], "get_media_requests") # first hook called
@ -315,7 +313,7 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase):
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, rsp1)])
self.assertEqual(new_item["results"], [(True, {})])
# rsp2 is ignored, rsp1 must be in results because request fingerprints are the same
req2 = Request(
@ -325,7 +323,7 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase):
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, rsp1)])
self.assertEqual(new_item["results"], [(True, {})])
@inlineCallbacks
def test_results_are_cached_for_requests_of_single_item(self):
@ -337,7 +335,7 @@ 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, rsp1), (True, rsp1)])
self.assertEqual(new_item["results"], [(True, {}), (True, {})])
@inlineCallbacks
def test_wait_if_request_is_downloading(self):
@ -363,7 +361,7 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase):
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, rsp1), (True, rsp1)])
self.assertEqual(new_item["results"], [(True, {}), (True, {})])
@inlineCallbacks
def test_use_media_to_download_result(self):
@ -376,57 +374,15 @@ class MediaPipelineTestCase(BaseMediaPipelineTestCase):
["get_media_requests", "media_to_download", "item_completed"],
)
class MockedMediaPipelineDeprecatedMethods(ImagesPipeline):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self._mockcalled = []
def get_media_requests(self, item, info):
item_url = item["image_urls"][0]
output_img = io.BytesIO()
img = Image.new("RGB", (60, 30), color="red")
img.save(output_img, format="JPEG")
return Request(
item_url,
meta={
"response": Response(item_url, status=200, body=output_img.getvalue())
},
def test_key_for_pipe(self):
self.assertEqual(
self.pipe._key_for_pipe("IMAGES", base_class_name="MediaPipeline"),
"MOCKEDMEDIAPIPELINE_IMAGES",
)
def inc_stats(self, *args, **kwargs):
return True
def media_to_download(self, request, info):
self._mockcalled.append("media_to_download")
return super().media_to_download(request, info)
def media_downloaded(self, response, request, info):
self._mockcalled.append("media_downloaded")
return super().media_downloaded(response, request, info)
def file_downloaded(self, response, request, info):
self._mockcalled.append("file_downloaded")
return super().file_downloaded(response, request, info)
def file_path(self, request, response=None, info=None):
self._mockcalled.append("file_path")
return super().file_path(request, response, info)
def thumb_path(self, request, thumb_id, response=None, info=None):
self._mockcalled.append("thumb_path")
return super().thumb_path(request, thumb_id, response, info)
def get_images(self, response, request, info):
self._mockcalled.append("get_images")
return super().get_images(response, request, info)
def image_downloaded(self, response, request, info):
self._mockcalled.append("image_downloaded")
return super().image_downloaded(response, request, info)
class MediaPipelineAllowRedirectSettingsTestCase(unittest.TestCase):
def _assert_request_no3xx(self, pipeline_class, settings):
pipe = pipeline_class(settings=Settings(settings))
request = Request("http://url")
@ -452,18 +408,11 @@ class MediaPipelineAllowRedirectSettingsTestCase(unittest.TestCase):
else:
self.assertNotIn(status, request.meta["handle_httpstatus_list"])
def test_standard_setting(self):
self._assert_request_no3xx(MediaPipeline, {"MEDIA_ALLOW_REDIRECTS": True})
def test_subclass_standard_setting(self):
class UserDefinedPipeline(MediaPipeline):
pass
self._assert_request_no3xx(UserDefinedPipeline, {"MEDIA_ALLOW_REDIRECTS": True})
def test_subclass_specific_setting(self):
class UserDefinedPipeline(MediaPipeline):
pass
self._assert_request_no3xx(
UserDefinedPipeline, {"USERDEFINEDPIPELINE_MEDIA_ALLOW_REDIRECTS": True}