#!/usr/bin/env python3 # -*- coding: utf-8 -*- import hashlib import logging import re import time import urllib.parse from collections.abc import Callable from typing import Optional import _hashlib import requests from praw.models import Submission from bdfr.exceptions import BulkDownloaderException logger = logging.getLogger(__name__) class Resource: def __init__(self, source_submission: Submission, url: str, download_function: Callable, extension: str = None): self.source_submission = source_submission self.content: Optional[bytes] = None self.url = url self.hash: Optional[_hashlib.HASH] = None self.extension = self._normalize_extension(extension) self.download_function = download_function if not self.extension: self.extension = self._determine_extension() @staticmethod def retry_download(url: str) -> Callable: return lambda global_params: Resource.http_download(url, global_params) def download(self, download_parameters: Optional[dict] = None): if download_parameters is None: download_parameters = {} if not self.content: try: content = self.download_function(download_parameters) except requests.exceptions.ConnectionError as e: raise BulkDownloaderException(f"Could not download resource: {e}") except BulkDownloaderException: raise if content: self.content = content # If we didn't have an extension before, try to detect from content if not self.extension: logger.debug(f"Attempting content-based extension detection for {self.url}") detected = self._detect_extension_by_content() self.extension = self._normalize_extension(detected) if detected else None if not self.hash and self.content: self.create_hash() def create_hash(self): self.hash = hashlib.md5(self.content) def _determine_extension(self) -> Optional[str]: extension_pattern = re.compile(r".*(\..{3,5})$") stripped_url = urllib.parse.urlsplit(self.url).path # Special handling for Reddit media URLs if self.url.startswith("https://www.reddit.com/media"): logger.debug(f"Detected Reddit media URL: {self.url}") parsed_url = urllib.parse.urlparse(self.url) url_param = urllib.parse.parse_qs(parsed_url.query).get('url', [None])[0] if url_param: decoded_url = urllib.parse.unquote(url_param) logger.debug(f"Reddit media URL decoded to: {decoded_url}") stripped_url = urllib.parse.urlsplit(decoded_url).path # Also handle preview.redd.it URLs which might not have extensions elif "preview.redd.it" in self.url and not stripped_url.endswith(('.jpg', '.jpeg', '.png', '.gif', '.webp')): logger.debug(f"Detected preview.redd.it URL without extension: {self.url}") # For preview URLs, try to infer from common patterns or add fallback logic match = re.search(extension_pattern, stripped_url) if match: extension = match.group(1) logger.debug(f"URL {self.url} -> extracted extension: {extension} (from path: {stripped_url})") return self._normalize_extension(extension) else: logger.warning(f"Could not determine extension for URL: {self.url} (path: {stripped_url})") # As a last resort, if we have content, try to detect by magic numbers if hasattr(self, 'content') and self.content: detected = self._detect_extension_by_content() return self._normalize_extension(detected) if detected else None return None def _detect_extension_by_content(self) -> Optional[str]: """Detect file extension by examining file content (magic numbers)""" if not self.content or len(self.content) < 16: return None # Check for common image formats if self.content.startswith(b'\xFF\xD8\xFF'): logger.debug(f"Detected JPEG by magic number for URL: {self.url}") return '.jpg' elif self.content.startswith(b'\x89PNG\r\n\x1a\n'): logger.debug(f"Detected PNG by magic number for URL: {self.url}") return '.png' elif self.content.startswith(b'GIF87a') or self.content.startswith(b'GIF89a'): logger.debug(f"Detected GIF by magic number for URL: {self.url}") return '.gif' elif self.content.startswith(b'RIFF') and self.content[8:12] == b'WEBP': logger.debug(f"Detected WebP by magic number for URL: {self.url}") return '.webp' elif self.content.startswith(b'BM'): logger.debug(f"Detected BMP by magic number for URL: {self.url}") return '.bmp' logger.debug(f"Could not detect file type by magic number for URL: {self.url}") return None def _normalize_extension(self, extension: Optional[str]) -> Optional[str]: """Normalize extension to lowercase for consistency""" if not extension: return None normalized = extension.lower() logger.debug(f"Normalized extension '{extension}' to '{normalized}'") return normalized @staticmethod def http_download(url: str, download_parameters: dict) -> Optional[bytes]: headers = download_parameters.get("headers") current_wait_time = 60 if "max_wait_time" in download_parameters: max_wait_time = download_parameters["max_wait_time"] else: max_wait_time = 300 while True: try: response = requests.get(url, headers=headers) if re.match(r"^2\d{2}", str(response.status_code)) and response.content: return response.content elif response.status_code in (408, 429): raise requests.exceptions.ConnectionError(f"Response code {response.status_code}") else: raise BulkDownloaderException( f"Unrecoverable error requesting resource: HTTP Code {response.status_code}" ) except (requests.exceptions.ConnectionError, requests.exceptions.ChunkedEncodingError) as e: logger.warning(f"Error occured downloading from {url}, waiting {current_wait_time} seconds: {e}") time.sleep(current_wait_time) if current_wait_time < max_wait_time: current_wait_time += 60 else: logger.error(f"Max wait time exceeded for resource at url {url}") raise