Full typing for scrapy/pipelines. (#6387)

This commit is contained in:
Andrey Rakhmatullin 2024-06-05 08:33:45 +04:00 committed by GitHub
parent 2b9e32f1ca
commit e56b425198
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
4 changed files with 424 additions and 143 deletions

View File

@ -10,6 +10,7 @@ from twisted.internet.defer import Deferred
from scrapy import Spider
from scrapy.middleware import MiddlewareManager
from scrapy.settings import Settings
from scrapy.utils.conf import build_component_list
from scrapy.utils.defer import deferred_f_from_coro_f
@ -18,7 +19,7 @@ class ItemPipelineManager(MiddlewareManager):
component_name = "item pipeline"
@classmethod
def _get_mwlist_from_settings(cls, settings) -> List[Any]:
def _get_mwlist_from_settings(cls, settings: Settings) -> List[Any]:
return build_component_list(settings.getwithbase("ITEM_PIPELINES"))
def _add_middleware(self, pipe: Any) -> None:

View File

@ -18,16 +18,35 @@ from ftplib import FTP
from io import BytesIO
from os import PathLike
from pathlib import Path
from typing import IO, TYPE_CHECKING, DefaultDict, Optional, Set, Type, Union, cast
from typing import (
IO,
TYPE_CHECKING,
Any,
Callable,
DefaultDict,
Dict,
List,
NoReturn,
Optional,
Protocol,
Set,
Type,
TypedDict,
Union,
cast,
)
from urllib.parse import urlparse
from itemadapter import ItemAdapter
from twisted.internet import defer, threads
from twisted.internet.defer import Deferred
from twisted.python.failure import Failure
from scrapy import Spider
from scrapy.exceptions import IgnoreRequest, NotConfigured
from scrapy.http import Request
from scrapy.http import Request, Response
from scrapy.http.request import NO_CALLBACK
from scrapy.pipelines.media import MediaPipeline
from scrapy.pipelines.media import FileInfo, FileInfoOrError, MediaPipeline
from scrapy.settings import Settings
from scrapy.utils.boto import is_botocore_available
from scrapy.utils.datatypes import CaseInsensitiveDict
@ -40,10 +59,11 @@ if TYPE_CHECKING:
# typing.Self requires Python 3.11
from typing_extensions import Self
logger = logging.getLogger(__name__)
def _to_string(path: Union[str, PathLike]) -> str:
def _to_string(path: Union[str, PathLike[str]]) -> str:
return str(path) # convert a Path object to string
@ -68,23 +88,54 @@ class FileException(Exception):
"""General media error exception"""
class StatInfo(TypedDict, total=False):
checksum: str
last_modified: float
class FilesStoreProtocol(Protocol):
def __init__(self, basedir: str): ...
def persist_file(
self,
path: str,
buf: BytesIO,
info: MediaPipeline.SpiderInfo,
meta: Optional[Dict[str, Any]] = None,
headers: Optional[Dict[str, str]] = None,
) -> Optional[Deferred[Any]]: ...
def stat_file(
self, path: str, info: MediaPipeline.SpiderInfo
) -> Union[StatInfo, Deferred[StatInfo]]: ...
class FSFilesStore:
def __init__(self, basedir: Union[str, PathLike]):
def __init__(self, basedir: Union[str, PathLike[str]]):
basedir = _to_string(basedir)
if "://" in basedir:
basedir = basedir.split("://", 1)[1]
self.basedir = basedir
self.basedir: str = basedir
self._mkdir(Path(self.basedir))
self.created_directories: DefaultDict[str, Set[str]] = defaultdict(set)
self.created_directories: DefaultDict[MediaPipeline.SpiderInfo, Set[str]] = (
defaultdict(set)
)
def persist_file(
self, path: Union[str, PathLike], buf, info, meta=None, headers=None
):
self,
path: Union[str, PathLike[str]],
buf: BytesIO,
info: MediaPipeline.SpiderInfo,
meta: Optional[Dict[str, Any]] = None,
headers: Optional[Dict[str, str]] = None,
) -> None:
absolute_path = self._get_filesystem_path(path)
self._mkdir(absolute_path.parent, info)
absolute_path.write_bytes(buf.getvalue())
def stat_file(self, path: Union[str, PathLike], info):
def stat_file(
self, path: Union[str, PathLike[str]], info: MediaPipeline.SpiderInfo
) -> StatInfo:
absolute_path = self._get_filesystem_path(path)
try:
last_modified = absolute_path.stat().st_mtime
@ -96,12 +147,14 @@ class FSFilesStore:
return {"last_modified": last_modified, "checksum": checksum}
def _get_filesystem_path(self, path: Union[str, PathLike]) -> Path:
def _get_filesystem_path(self, path: Union[str, PathLike[str]]) -> Path:
path_comps = _to_string(path).split("/")
return Path(self.basedir, *path_comps)
def _mkdir(self, dirname: Path, domain: Optional[str] = None):
seen = self.created_directories[domain] if domain else set()
def _mkdir(
self, dirname: Path, domain: Optional[MediaPipeline.SpiderInfo] = None
) -> None:
seen: Set[str] = self.created_directories[domain] if domain else set()
if str(dirname) not in seen:
if not dirname.exists():
dirname.mkdir(parents=True)
@ -122,7 +175,7 @@ class S3FilesStore:
"Cache-Control": "max-age=172800",
}
def __init__(self, uri):
def __init__(self, uri: str):
if not is_botocore_available():
raise NotConfigured("missing botocore library")
import botocore.session
@ -142,8 +195,10 @@ class S3FilesStore:
raise ValueError(f"Incorrect URI scheme in {uri}, expected 's3'")
self.bucket, self.prefix = uri[5:].split("/", 1)
def stat_file(self, path, info):
def _onsuccess(boto_key):
def stat_file(
self, path: str, info: MediaPipeline.SpiderInfo
) -> Deferred[StatInfo]:
def _onsuccess(boto_key: Dict[str, Any]) -> StatInfo:
checksum = boto_key["ETag"].strip('"')
last_modified = boto_key["LastModified"]
modified_stamp = time.mktime(last_modified.timetuple())
@ -151,13 +206,23 @@ class S3FilesStore:
return self._get_boto_key(path).addCallback(_onsuccess)
def _get_boto_key(self, path):
def _get_boto_key(self, path: str) -> Deferred[Dict[str, Any]]:
key_name = f"{self.prefix}{path}"
return threads.deferToThread(
self.s3_client.head_object, Bucket=self.bucket, Key=key_name
return cast(
"Deferred[Dict[str, Any]]",
threads.deferToThread(
self.s3_client.head_object, Bucket=self.bucket, Key=key_name # type: ignore[attr-defined]
),
)
def persist_file(self, path, buf, info, meta=None, headers=None):
def persist_file(
self,
path: str,
buf: BytesIO,
info: MediaPipeline.SpiderInfo,
meta: Optional[Dict[str, Any]] = None,
headers: Optional[Dict[str, str]] = None,
) -> Deferred[Any]:
"""Upload file to S3 storage"""
key_name = f"{self.prefix}{path}"
buf.seek(0)
@ -165,7 +230,7 @@ class S3FilesStore:
if headers:
extra.update(self._headers_to_botocore_kwargs(headers))
return threads.deferToThread(
self.s3_client.put_object,
self.s3_client.put_object, # type: ignore[attr-defined]
Bucket=self.bucket,
Key=key_name,
Body=buf,
@ -174,7 +239,7 @@ class S3FilesStore:
**extra,
)
def _headers_to_botocore_kwargs(self, headers):
def _headers_to_botocore_kwargs(self, headers: Dict[str, Any]) -> Dict[str, Any]:
"""Convert headers to botocore keyword arguments."""
# This is required while we need to support both boto and botocore.
mapping = CaseInsensitiveDict(
@ -206,7 +271,7 @@ class S3FilesStore:
"X-Amz-Website-Redirect-Location": "WebsiteRedirectLocation",
}
)
extra = {}
extra: Dict[str, Any] = {}
for key, value in headers.items():
try:
kwarg = mapping[key]
@ -226,13 +291,13 @@ class GCSFilesStore:
# Overridden from settings.FILES_STORE_GCS_ACL in FilesPipeline.from_settings.
POLICY = None
def __init__(self, uri):
def __init__(self, uri: str):
from google.cloud import storage
client = storage.Client(project=self.GCS_PROJECT_ID)
bucket, prefix = uri[5:].split("/", 1)
self.bucket = client.bucket(bucket)
self.prefix = prefix
self.prefix: str = prefix
permissions = self.bucket.test_iam_permissions(
["storage.objects.get", "storage.objects.create"]
)
@ -248,8 +313,10 @@ class GCSFilesStore:
{"bucket": bucket},
)
def stat_file(self, path, info):
def _onsuccess(blob):
def stat_file(
self, path: str, info: MediaPipeline.SpiderInfo
) -> Deferred[StatInfo]:
def _onsuccess(blob) -> StatInfo:
if blob:
checksum = base64.b64decode(blob.md5_hash).hex()
last_modified = time.mktime(blob.updated.timetuple())
@ -257,19 +324,29 @@ class GCSFilesStore:
return {}
blob_path = self._get_blob_path(path)
return threads.deferToThread(self.bucket.get_blob, blob_path).addCallback(
_onsuccess
return cast(
Deferred[StatInfo],
threads.deferToThread(self.bucket.get_blob, blob_path).addCallback(
_onsuccess
),
)
def _get_content_type(self, headers):
def _get_content_type(self, headers: Optional[Dict[str, str]]) -> str:
if headers and "Content-Type" in headers:
return headers["Content-Type"]
return "application/octet-stream"
def _get_blob_path(self, path):
def _get_blob_path(self, path: str) -> str:
return self.prefix + path
def persist_file(self, path, buf, info, meta=None, headers=None):
def persist_file(
self,
path: str,
buf: BytesIO,
info: MediaPipeline.SpiderInfo,
meta: Optional[Dict[str, Any]] = None,
headers: Optional[Dict[str, str]] = None,
) -> Deferred[Any]:
blob_path = self._get_blob_path(path)
blob = self.bucket.blob(blob_path)
blob.cache_control = self.CACHE_CONTROL
@ -283,22 +360,33 @@ class GCSFilesStore:
class FTPFilesStore:
FTP_USERNAME = None
FTP_PASSWORD = None
USE_ACTIVE_MODE = None
FTP_USERNAME: Optional[str] = None
FTP_PASSWORD: Optional[str] = None
USE_ACTIVE_MODE: Optional[bool] = None
def __init__(self, uri):
def __init__(self, uri: str):
if not uri.startswith("ftp://"):
raise ValueError(f"Incorrect URI scheme in {uri}, expected 'ftp'")
u = urlparse(uri)
self.port = u.port
self.host = u.hostname
assert u.port
assert u.hostname
self.port: int = u.port
self.host: str = u.hostname
self.port = int(u.port or 21)
self.username = u.username or self.FTP_USERNAME
self.password = u.password or self.FTP_PASSWORD
self.basedir = u.path.rstrip("/")
assert self.FTP_USERNAME
assert self.FTP_PASSWORD
self.username: str = u.username or self.FTP_USERNAME
self.password: str = u.password or self.FTP_PASSWORD
self.basedir: str = u.path.rstrip("/")
def persist_file(self, path, buf, info, meta=None, headers=None):
def persist_file(
self,
path: str,
buf: BytesIO,
info: MediaPipeline.SpiderInfo,
meta: Optional[Dict[str, Any]] = None,
headers: Optional[Dict[str, str]] = None,
) -> Deferred[Any]:
path = f"{self.basedir}/{path}"
return threads.deferToThread(
ftp_store_file,
@ -311,8 +399,10 @@ class FTPFilesStore:
use_active_mode=self.USE_ACTIVE_MODE,
)
def stat_file(self, path, info):
def _stat_file(path):
def stat_file(
self, path: str, info: MediaPipeline.SpiderInfo
) -> Deferred[StatInfo]:
def _stat_file(path: str) -> StatInfo:
try:
ftp = FTP()
ftp.connect(self.host, self.port)
@ -328,7 +418,7 @@ class FTPFilesStore:
except Exception:
return {}
return threads.deferToThread(_stat_file, path)
return cast("Deferred[StatInfo]", threads.deferToThread(_stat_file, path))
class FilesPipeline(MediaPipeline):
@ -350,20 +440,23 @@ class FilesPipeline(MediaPipeline):
"""
MEDIA_NAME = "file"
EXPIRES = 90
STORE_SCHEMES = {
MEDIA_NAME: str = "file"
EXPIRES: int = 90
STORE_SCHEMES: Dict[str, Type[FilesStoreProtocol]] = {
"": FSFilesStore,
"file": FSFilesStore,
"s3": S3FilesStore,
"gs": GCSFilesStore,
"ftp": FTPFilesStore,
}
DEFAULT_FILES_URLS_FIELD = "file_urls"
DEFAULT_FILES_RESULT_FIELD = "files"
DEFAULT_FILES_URLS_FIELD: str = "file_urls"
DEFAULT_FILES_RESULT_FIELD: str = "files"
def __init__(
self, store_uri: Union[str, PathLike], download_func=None, settings=None
self,
store_uri: Union[str, PathLike[str]],
download_func: Optional[Callable[[Request, Spider], Response]] = None,
settings: Union[Settings, Dict[str, Any], None] = None,
):
store_uri = _to_string(store_uri)
if not store_uri:
@ -372,26 +465,26 @@ class FilesPipeline(MediaPipeline):
if isinstance(settings, dict) or settings is None:
settings = Settings(settings)
cls_name = "FilesPipeline"
self.store = self._get_store(store_uri)
self.store: FilesStoreProtocol = self._get_store(store_uri)
resolve = functools.partial(
self._key_for_pipe, base_class_name=cls_name, settings=settings
)
self.expires = settings.getint(resolve("FILES_EXPIRES"), self.EXPIRES)
self.expires: int = settings.getint(resolve("FILES_EXPIRES"), self.EXPIRES)
if not hasattr(self, "FILES_URLS_FIELD"):
self.FILES_URLS_FIELD = self.DEFAULT_FILES_URLS_FIELD
if not hasattr(self, "FILES_RESULT_FIELD"):
self.FILES_RESULT_FIELD = self.DEFAULT_FILES_RESULT_FIELD
self.files_urls_field = settings.get(
self.files_urls_field: str = settings.get(
resolve("FILES_URLS_FIELD"), self.FILES_URLS_FIELD
)
self.files_result_field = settings.get(
self.files_result_field: str = settings.get(
resolve("FILES_RESULT_FIELD"), self.FILES_RESULT_FIELD
)
super().__init__(download_func=download_func, settings=settings)
@classmethod
def from_settings(cls, settings) -> Self:
def from_settings(cls, settings: Settings) -> Self:
s3store: Type[S3FilesStore] = cast(Type[S3FilesStore], cls.STORE_SCHEMES["s3"])
s3store.AWS_ACCESS_KEY_ID = settings["AWS_ACCESS_KEY_ID"]
s3store.AWS_SECRET_ACCESS_KEY = settings["AWS_SECRET_ACCESS_KEY"]
@ -418,7 +511,7 @@ class FilesPipeline(MediaPipeline):
store_uri = settings["FILES_STORE"]
return cls(store_uri, settings=settings)
def _get_store(self, uri: str):
def _get_store(self, uri: str) -> FilesStoreProtocol:
if Path(uri).is_absolute(): # to support win32 paths like: C:\\some\dir
scheme = "file"
else:
@ -426,19 +519,21 @@ class FilesPipeline(MediaPipeline):
store_cls = self.STORE_SCHEMES[scheme]
return store_cls(uri)
def media_to_download(self, request, info, *, item=None):
def _onsuccess(result):
def media_to_download(
self, request: Request, info: MediaPipeline.SpiderInfo, *, item: Any = None
) -> Deferred[Optional[FileInfo]]:
def _onsuccess(result: StatInfo) -> Optional[FileInfo]:
if not result:
return # returning None force download
return None # returning None force download
last_modified = result.get("last_modified", None)
if not last_modified:
return # returning None force download
return None # returning None force download
age_seconds = time.time() - last_modified
age_days = age_seconds / 60 / 60 / 24
if age_days > self.expires:
return # returning None force download
return None # returning None force download
referer = referer_str(request)
logger.debug(
@ -458,19 +553,22 @@ class FilesPipeline(MediaPipeline):
}
path = self.file_path(request, info=info, item=item)
dfd = defer.maybeDeferred(self.store.stat_file, path, info)
dfd.addCallback(_onsuccess)
dfd.addErrback(lambda _: None)
dfd.addErrback(
# defer.maybeDeferred() overloads don't seem to support a Union[_T, Deferred[_T]] return type
dfd: Deferred[StatInfo] = defer.maybeDeferred(self.store.stat_file, path, info) # type: ignore[arg-type]
dfd2: Deferred[Optional[FileInfo]] = dfd.addCallback(_onsuccess)
dfd2.addErrback(lambda _: None)
dfd2.addErrback(
lambda f: logger.error(
self.__class__.__name__ + ".store.stat_file",
exc_info=failure_to_exc_info(f),
extra={"spider": info.spider},
)
)
return dfd
return dfd2
def media_failed(self, failure, request, info):
def media_failed(
self, failure: Failure, request: Request, info: MediaPipeline.SpiderInfo
) -> NoReturn:
if not isinstance(failure.value, IgnoreRequest):
referer = referer_str(request)
logger.warning(
@ -487,7 +585,14 @@ class FilesPipeline(MediaPipeline):
raise FileException
def media_downloaded(self, response, request, info, *, item=None):
def media_downloaded(
self,
response: Response,
request: Request,
info: MediaPipeline.SpiderInfo,
*,
item: Any = None,
) -> FileInfo:
referer = referer_str(request)
if response.status != 200:
@ -546,16 +651,26 @@ class FilesPipeline(MediaPipeline):
"status": status,
}
def inc_stats(self, spider, status):
def inc_stats(self, spider: Spider, status: str) -> None:
assert spider.crawler.stats
spider.crawler.stats.inc_value("file_count", spider=spider)
spider.crawler.stats.inc_value(f"file_status_count/{status}", spider=spider)
# Overridable Interface
def get_media_requests(self, item, info):
def get_media_requests(
self, item: Any, info: MediaPipeline.SpiderInfo
) -> List[Request]:
urls = ItemAdapter(item).get(self.files_urls_field, [])
return [Request(u, callback=NO_CALLBACK) for u in urls]
def file_downloaded(self, response, request, info, *, item=None):
def file_downloaded(
self,
response: Response,
request: Request,
info: MediaPipeline.SpiderInfo,
*,
item: Any = None,
) -> str:
path = self.file_path(request, response=response, info=info, item=item)
buf = BytesIO(response.body)
checksum = _md5sum(buf)
@ -563,12 +678,21 @@ class FilesPipeline(MediaPipeline):
self.store.persist_file(path, buf, info)
return checksum
def item_completed(self, results, item, info):
def item_completed(
self, results: List[FileInfoOrError], item: Any, info: MediaPipeline.SpiderInfo
) -> Any:
with suppress(KeyError):
ItemAdapter(item)[self.files_result_field] = [x for ok, x in results if ok]
return item
def file_path(self, request, response=None, info=None, *, item=None):
def file_path(
self,
request: Request,
response: Optional[Response] = None,
info: Optional[MediaPipeline.SpiderInfo] = None,
*,
item: Any = None,
) -> str:
media_guid = hashlib.sha1(to_bytes(request.url)).hexdigest() # nosec
media_ext = Path(request.url).suffix
# Handles empty and wild extensions by trying to guess the
@ -577,5 +701,5 @@ class FilesPipeline(MediaPipeline):
media_ext = ""
media_type = mimetypes.guess_type(request.url)[0]
if media_type:
media_ext = mimetypes.guess_extension(media_type)
media_ext = cast(str, mimetypes.guess_extension(media_type))
return f"full/{media_guid}{media_ext}"

View File

@ -12,12 +12,25 @@ import warnings
from contextlib import suppress
from io import BytesIO
from os import PathLike
from typing import TYPE_CHECKING, Dict, Tuple, Type, Union, cast
from typing import (
TYPE_CHECKING,
Any,
Callable,
Dict,
Iterable,
List,
Optional,
Tuple,
Type,
Union,
cast,
)
from itemadapter import ItemAdapter
from scrapy import Spider
from scrapy.exceptions import DropItem, NotConfigured, ScrapyDeprecationWarning
from scrapy.http import Request
from scrapy.http import Request, Response
from scrapy.http.request import NO_CALLBACK
from scrapy.pipelines.files import (
FileException,
@ -27,20 +40,20 @@ from scrapy.pipelines.files import (
S3FilesStore,
_md5sum,
)
# TODO: from scrapy.pipelines.media import MediaPipeline
from scrapy.pipelines.media import FileInfoOrError, MediaPipeline
from scrapy.settings import Settings
from scrapy.utils.python import get_func_args, to_bytes
if TYPE_CHECKING:
# typing.Self requires Python 3.11
from PIL import Image
from typing_extensions import Self
class NoimagesDrop(DropItem):
"""Product with no images exception"""
def __init__(self, *args, **kwargs):
def __init__(self, *args: Any, **kwargs: Any):
warnings.warn(
"The NoimagesDrop class is deprecated",
category=ScrapyDeprecationWarning,
@ -56,19 +69,22 @@ class ImageException(FileException):
class ImagesPipeline(FilesPipeline):
"""Abstract pipeline that implement the image thumbnail generation logic"""
MEDIA_NAME = "image"
MEDIA_NAME: str = "image"
# Uppercase attributes kept for backward compatibility with code that subclasses
# ImagesPipeline. They may be overridden by settings.
MIN_WIDTH = 0
MIN_HEIGHT = 0
EXPIRES = 90
MIN_WIDTH: int = 0
MIN_HEIGHT: int = 0
EXPIRES: int = 90
THUMBS: Dict[str, Tuple[int, int]] = {}
DEFAULT_IMAGES_URLS_FIELD = "image_urls"
DEFAULT_IMAGES_RESULT_FIELD = "images"
def __init__(
self, store_uri: Union[str, PathLike], download_func=None, settings=None
self,
store_uri: Union[str, PathLike[str]],
download_func: Optional[Callable[[Request, Spider], Response]] = None,
settings: Union[Settings, Dict[str, Any], None] = None,
):
try:
from PIL import Image
@ -89,27 +105,33 @@ class ImagesPipeline(FilesPipeline):
base_class_name="ImagesPipeline",
settings=settings,
)
self.expires = settings.getint(resolve("IMAGES_EXPIRES"), self.EXPIRES)
self.expires: int = settings.getint(resolve("IMAGES_EXPIRES"), self.EXPIRES)
if not hasattr(self, "IMAGES_RESULT_FIELD"):
self.IMAGES_RESULT_FIELD = self.DEFAULT_IMAGES_RESULT_FIELD
self.IMAGES_RESULT_FIELD: str = self.DEFAULT_IMAGES_RESULT_FIELD
if not hasattr(self, "IMAGES_URLS_FIELD"):
self.IMAGES_URLS_FIELD = self.DEFAULT_IMAGES_URLS_FIELD
self.IMAGES_URLS_FIELD: str = self.DEFAULT_IMAGES_URLS_FIELD
self.images_urls_field = settings.get(
self.images_urls_field: str = settings.get(
resolve("IMAGES_URLS_FIELD"), self.IMAGES_URLS_FIELD
)
self.images_result_field = settings.get(
self.images_result_field: str = settings.get(
resolve("IMAGES_RESULT_FIELD"), self.IMAGES_RESULT_FIELD
)
self.min_width = settings.getint(resolve("IMAGES_MIN_WIDTH"), self.MIN_WIDTH)
self.min_height = settings.getint(resolve("IMAGES_MIN_HEIGHT"), self.MIN_HEIGHT)
self.thumbs = settings.get(resolve("IMAGES_THUMBS"), self.THUMBS)
self.min_width: int = settings.getint(
resolve("IMAGES_MIN_WIDTH"), self.MIN_WIDTH
)
self.min_height: int = settings.getint(
resolve("IMAGES_MIN_HEIGHT"), self.MIN_HEIGHT
)
self.thumbs: Dict[str, Tuple[int, int]] = settings.get(
resolve("IMAGES_THUMBS"), self.THUMBS
)
self._deprecated_convert_image = None
self._deprecated_convert_image: Optional[bool] = None
@classmethod
def from_settings(cls, settings) -> Self:
def from_settings(cls, settings: Settings) -> Self:
s3store: Type[S3FilesStore] = cast(Type[S3FilesStore], cls.STORE_SCHEMES["s3"])
s3store.AWS_ACCESS_KEY_ID = settings["AWS_ACCESS_KEY_ID"]
s3store.AWS_SECRET_ACCESS_KEY = settings["AWS_SECRET_ACCESS_KEY"]
@ -136,11 +158,25 @@ class ImagesPipeline(FilesPipeline):
store_uri = settings["IMAGES_STORE"]
return cls(store_uri, settings=settings)
def file_downloaded(self, response, request, info, *, item=None):
def file_downloaded(
self,
response: Response,
request: Request,
info: MediaPipeline.SpiderInfo,
*,
item: Any = None,
) -> str:
return self.image_downloaded(response, request, info, item=item)
def image_downloaded(self, response, request, info, *, item=None):
checksum = None
def image_downloaded(
self,
response: Response,
request: Request,
info: MediaPipeline.SpiderInfo,
*,
item: Any = None,
) -> str:
checksum: Optional[str] = None
for path, image, buf in self.get_images(response, request, info, item=item):
if checksum is None:
buf.seek(0)
@ -153,9 +189,17 @@ class ImagesPipeline(FilesPipeline):
meta={"width": width, "height": height},
headers={"Content-Type": "image/jpeg"},
)
assert checksum is not None
return checksum
def get_images(self, response, request, info, *, item=None):
def get_images(
self,
response: Response,
request: Request,
info: MediaPipeline.SpiderInfo,
*,
item: Any = None,
) -> Iterable[Tuple[str, Image.Image, BytesIO]]:
path = self.file_path(request, response=response, info=info, item=item)
orig_image = self._Image.open(BytesIO(response.body))
@ -196,7 +240,12 @@ class ImagesPipeline(FilesPipeline):
thumb_image, thumb_buf = self.convert_image(image, size, buf)
yield thumb_path, thumb_image, thumb_buf
def convert_image(self, image, size=None, response_body=None):
def convert_image(
self,
image: Image.Image,
size: Optional[Tuple[int, int]] = None,
response_body: Optional[BytesIO] = None,
) -> Tuple[Image.Image, BytesIO]:
if response_body is None:
warnings.warn(
f"{self.__class__.__name__}.convert_image() method called in a deprecated way, "
@ -225,7 +274,7 @@ class ImagesPipeline(FilesPipeline):
# when updating the minimum requirements for Pillow.
resampling_filter = self._Image.Resampling.LANCZOS
except AttributeError:
resampling_filter = self._Image.ANTIALIAS
resampling_filter = self._Image.ANTIALIAS # type: ignore[attr-defined]
image.thumbnail(size, resampling_filter)
elif response_body is not None and image.format == "JPEG":
return image, response_body
@ -234,19 +283,38 @@ class ImagesPipeline(FilesPipeline):
image.save(buf, "JPEG")
return image, buf
def get_media_requests(self, item, info):
def get_media_requests(
self, item: Any, info: MediaPipeline.SpiderInfo
) -> List[Request]:
urls = ItemAdapter(item).get(self.images_urls_field, [])
return [Request(u, callback=NO_CALLBACK) for u in urls]
def item_completed(self, results, item, info):
def item_completed(
self, results: List[FileInfoOrError], item: Any, info: MediaPipeline.SpiderInfo
) -> Any:
with suppress(KeyError):
ItemAdapter(item)[self.images_result_field] = [x for ok, x in results if ok]
return item
def file_path(self, request, response=None, info=None, *, item=None):
def file_path(
self,
request: Request,
response: Optional[Response] = None,
info: Optional[MediaPipeline.SpiderInfo] = None,
*,
item: Any = None,
) -> str:
image_guid = hashlib.sha1(to_bytes(request.url)).hexdigest() # nosec
return f"full/{image_guid}.jpg"
def thumb_path(self, request, thumb_id, response=None, info=None, *, item=None):
def thumb_path(
self,
request: Request,
thumb_id: str,
response: Optional[Response] = None,
info: Optional[MediaPipeline.SpiderInfo] = None,
*,
item: Any = None,
) -> str:
thumb_guid = hashlib.sha1(to_bytes(request.url)).hexdigest() # nosec
return f"thumbs/{thumb_id}/{thumb_guid}.jpg"

View File

@ -4,54 +4,101 @@ import functools
import logging
from abc import ABC, abstractmethod
from collections import defaultdict
from typing import TYPE_CHECKING
from typing import (
TYPE_CHECKING,
Any,
Callable,
DefaultDict,
Dict,
List,
Literal,
NoReturn,
Optional,
Set,
Tuple,
TypedDict,
TypeVar,
Union,
cast,
)
from twisted.internet.defer import Deferred, DeferredList
from twisted.python.failure import Failure
from scrapy.http.request import NO_CALLBACK
from scrapy import Spider
from scrapy.crawler import Crawler
from scrapy.http import Response
from scrapy.http.request import NO_CALLBACK, Request
from scrapy.settings import Settings
from scrapy.utils.datatypes import SequenceExclude
from scrapy.utils.defer import defer_result, mustbe_deferred
from scrapy.utils.log import failure_to_exc_info
from scrapy.utils.misc import arg_to_iter
from scrapy.utils.request import RequestFingerprinter
if TYPE_CHECKING:
# typing.Self requires Python 3.11
from typing_extensions import Self
_T = TypeVar("_T")
class FileInfo(TypedDict):
url: str
path: str
checksum: Optional[str]
status: str
FileInfoOrError = Union[Tuple[Literal[True], FileInfo], Tuple[Literal[False], Failure]]
logger = logging.getLogger(__name__)
class MediaPipeline(ABC):
LOG_FAILED_RESULTS = True
crawler: Crawler
_fingerprinter: RequestFingerprinter
LOG_FAILED_RESULTS: bool = True
class SpiderInfo:
def __init__(self, spider):
self.spider = spider
self.downloading = set()
self.downloaded = {}
self.waiting = defaultdict(list)
def __init__(self, spider: Spider):
self.spider: Spider = spider
self.downloading: Set[bytes] = set()
self.downloaded: Dict[bytes, Union[FileInfo, Failure]] = {}
self.waiting: DefaultDict[bytes, List[Deferred[FileInfo]]] = defaultdict(
list
)
def __init__(self, download_func=None, settings=None):
def __init__(
self,
download_func: Optional[Callable[[Request, Spider], Response]] = None,
settings: Union[Settings, Dict[str, Any], None] = None,
):
self.download_func = download_func
self._expects_item = {}
if isinstance(settings, dict) or settings is None:
settings = Settings(settings)
resolve = functools.partial(
self._key_for_pipe, base_class_name="MediaPipeline", settings=settings
)
self.allow_redirects = settings.getbool(resolve("MEDIA_ALLOW_REDIRECTS"), False)
self.allow_redirects: bool = settings.getbool(
resolve("MEDIA_ALLOW_REDIRECTS"), False
)
self._handle_statuses(self.allow_redirects)
def _handle_statuses(self, allow_redirects):
def _handle_statuses(self, allow_redirects: bool) -> None:
self.handle_httpstatus_list = None
if allow_redirects:
self.handle_httpstatus_list = SequenceExclude(range(300, 400))
def _key_for_pipe(self, key, base_class_name=None, settings=None):
def _key_for_pipe(
self,
key: str,
base_class_name: Optional[str] = None,
settings: Optional[Settings] = None,
) -> str:
class_name = self.__class__.__name__
formatted_key = f"{class_name.upper()}_{key}"
if (
@ -64,26 +111,34 @@ class MediaPipeline(ABC):
return formatted_key
@classmethod
def from_crawler(cls, crawler) -> Self:
def from_crawler(cls, crawler: Crawler) -> Self:
pipe: Self
try:
pipe = cls.from_settings(crawler.settings) # type: ignore[attr-defined]
except AttributeError:
pipe = cls()
pipe.crawler = crawler
assert crawler.request_fingerprinter
pipe._fingerprinter = crawler.request_fingerprinter
return pipe
def open_spider(self, spider):
def open_spider(self, spider: Spider) -> None:
self.spiderinfo = self.SpiderInfo(spider)
def process_item(self, item, spider):
def process_item(
self, item: Any, spider: Spider
) -> Deferred[List[FileInfoOrError]]:
info = self.spiderinfo
requests = arg_to_iter(self.get_media_requests(item, info))
dlist = [self._process_request(r, info, item) for r in requests]
dfd = DeferredList(dlist, consumeErrors=True)
dfd = cast(
"Deferred[List[FileInfoOrError]]", DeferredList(dlist, consumeErrors=True)
)
return dfd.addCallback(self.item_completed, item, info)
def _process_request(self, request, info, item):
def _process_request(
self, request: Request, info: SpiderInfo, item: Any
) -> Deferred[FileInfo]:
fp = self._fingerprinter.fingerprint(request)
eb = request.errback
request.callback = NO_CALLBACK
@ -97,7 +152,7 @@ class MediaPipeline(ABC):
return d
# Otherwise, wait for result
wad = Deferred()
wad: Deferred[FileInfo] = Deferred()
if eb:
wad.addErrback(eb)
info.waiting[fp].append(wad)
@ -108,36 +163,48 @@ class MediaPipeline(ABC):
# Download request checking media_to_download hook output first
info.downloading.add(fp)
dfd = mustbe_deferred(self.media_to_download, request, info, item=item)
dfd.addCallback(self._check_media_to_download, request, info, item=item)
dfd.addErrback(self._log_exception)
dfd.addBoth(self._cache_result_and_execute_waiters, fp, info)
return dfd.addBoth(lambda _: wad) # it must return wad at last
dfd: Deferred[Optional[FileInfo]] = mustbe_deferred(
self.media_to_download, request, info, item=item
)
dfd2: Deferred[FileInfo] = dfd.addCallback(
self._check_media_to_download, request, info, item=item
)
dfd2.addErrback(self._log_exception)
dfd2.addBoth(self._cache_result_and_execute_waiters, fp, info)
return dfd2.addBoth(lambda _: wad) # it must return wad at last
def _log_exception(self, result):
def _log_exception(self, result: Failure) -> Failure:
logger.exception(result)
return result
def _modify_media_request(self, request):
def _modify_media_request(self, request: Request) -> None:
if self.handle_httpstatus_list:
request.meta["handle_httpstatus_list"] = self.handle_httpstatus_list
else:
request.meta["handle_httpstatus_all"] = True
def _check_media_to_download(self, result, request, info, item):
def _check_media_to_download(
self, result: Optional[FileInfo], request: Request, info: SpiderInfo, item: Any
) -> Union[FileInfo, Deferred[FileInfo]]:
if result is not None:
return result
dfd: Deferred[Response]
if self.download_func:
# this ugly code was left only to support tests. TODO: remove
dfd = mustbe_deferred(self.download_func, request, info.spider)
else:
self._modify_media_request(request)
assert self.crawler.engine
dfd = self.crawler.engine.download(request)
dfd.addCallback(self.media_downloaded, request, info, item=item)
dfd.addErrback(self.media_failed, request, info)
return dfd
dfd2: Deferred[FileInfo] = dfd.addCallback(
self.media_downloaded, request, info, item=item
)
dfd2.addErrback(self.media_failed, request, info)
return dfd2
def _cache_result_and_execute_waiters(self, result, fp, info):
def _cache_result_and_execute_waiters(
self, result: Union[FileInfo, Failure], fp: bytes, info: SpiderInfo
) -> None:
if isinstance(result, Failure):
# minimize cached information for failure
result.cleanFailure()
@ -176,30 +243,44 @@ class MediaPipeline(ABC):
# Overridable Interface
@abstractmethod
def media_to_download(self, request, info, *, item=None):
def media_to_download(
self, request: Request, info: SpiderInfo, *, item: Any = None
) -> Deferred[Optional[FileInfo]]:
"""Check request before starting download"""
raise NotImplementedError()
@abstractmethod
def get_media_requests(self, item, info):
def get_media_requests(self, item: Any, info: SpiderInfo) -> List[Request]:
"""Returns the media requests to download"""
raise NotImplementedError()
@abstractmethod
def media_downloaded(self, response, request, info, *, item=None):
def media_downloaded(
self,
response: Response,
request: Request,
info: SpiderInfo,
*,
item: Any = None,
) -> FileInfo:
"""Handler for success downloads"""
raise NotImplementedError()
@abstractmethod
def media_failed(self, failure, request, info):
def media_failed(
self, failure: Failure, request: Request, info: SpiderInfo
) -> NoReturn:
"""Handler for failed downloads"""
raise NotImplementedError()
def item_completed(self, results, item, info):
def item_completed(
self, results: List[FileInfoOrError], item: Any, info: SpiderInfo
) -> Any:
"""Called per item when all media requests has been processed"""
if self.LOG_FAILED_RESULTS:
for ok, value in results:
if not ok:
assert isinstance(value, Failure)
logger.error(
"%(class)s found errors processing %(item)s",
{"class": self.__class__.__name__, "item": item},
@ -209,6 +290,13 @@ class MediaPipeline(ABC):
return item
@abstractmethod
def file_path(self, request, response=None, info=None, *, item=None):
def file_path(
self,
request: Request,
response: Optional[Response] = None,
info: Optional[SpiderInfo] = None,
*,
item: Any = None,
) -> str:
"""Returns the path where downloaded media should be stored"""
raise NotImplementedError()