mirror of https://github.com/scrapy/scrapy.git
Fix overridable methods in MediaPipeline (#6368)
This commit is contained in:
parent
986d1ee1dd
commit
cadb0dd707
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
Loading…
Reference in New Issue