Files
BDFR_Web/bdfr/resource.py
T
ModerateWinGuy 0157f462cc
formatting_check / formatting_check (push) Failing after 3s
Python Test / test (.ps1, windows-latest, 3.9) (push) Has been cancelled
Python Test / test (.sh, macos-latest, 3.9) (push) Has been cancelled
Python Test / test (.sh, ubuntu-latest, 3.9) (push) Has been cancelled
reformatting
2026-07-14 21:32:55 +12:00

170 lines
7.2 KiB
Python

#!/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
# Log the original extension before normalization
if extension:
logger.debug(f"Resource constructor received extension: '{extension}' for URL: {url}")
self.extension = self._normalize_extension(extension)
self.download_function = download_function
if not self.extension:
self.extension = self._determine_extension()
logger.debug(f"Extension determined from URL: '{self.extension}' for URL: {url}")
@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
original = extension
# Ensure extension starts with a dot
if not extension.startswith("."):
extension = "." + extension
normalized = extension.lower()
if original != normalized:
logger.info(
f"Extension normalization: '{original}' -> '{normalized}' for URL: {self.url if hasattr(self, 'url') else 'unknown'}"
)
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