Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion reddit2telegram/supplier.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,8 @@ def send_to_channel_from_subreddit(how_to_post, channel_to_post, subreddit, subm
client_id=config['reddit']['client_id'],
client_secret=config['reddit']['client_secret'],
username=config['reddit']['username'],
password=config['reddit']['password']
password=config['reddit']['password'],
requestor_kwargs={'timeout': utils.REDDIT_TIMEOUT}
)
if submissions_ranking == 'top':
submissions = reddit.subreddit(subreddit).top(limit=submissions_limit)
Expand Down
26 changes: 26 additions & 0 deletions reddit2telegram/task_queue.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,8 @@

TASK_SUPPLY = 'supply'

ABANDONED_TASK_MIN_AGE_SECONDS = 60


logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -67,6 +69,29 @@ def select_all_available_tasks(collection: Collection):
}, sort=[('created_at', pymongo.ASCENDING)]))


def recover_abandoned_tasks(collection: Collection):
cutoff = time.time() - ABANDONED_TASK_MIN_AGE_SECONDS
result = collection.update_many(
{
'status': {
'$in': [
TaskStatus.IN_PROGRESS.value,
TaskStatus.SCHEDULED.value,
]
},
'updated_at': {'$lt': cutoff},
},
{
'$set': {
'status': TaskStatus.NEW.value,
'updated_at': time.time(),
},
}
)
if result.modified_count:
logger.warning('Recovered %d abandoned tasks', result.modified_count)


def execute_task(collection: Collection, id: ObjectId, name: str, args: Mapping):
update_task_status(collection, id, TaskStatus.IN_PROGRESS)
try:
Expand All @@ -90,6 +115,7 @@ def start_consumer(
):
running = True
collection = mongo_database[COLLECTION]
recover_abandoned_tasks(collection)
logger.info('Starting consumer %s', executor)
while running:
try:
Expand Down
120 changes: 79 additions & 41 deletions reddit2telegram/utils/__init__.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,9 @@
# encoding:utf-8

import urllib
from urllib.parse import urlparse
from html import unescape
import requests
from requests.exceptions import InvalidSchema, MissingSchema
from requests.exceptions import InvalidSchema, MissingSchema, RequestException
import os
import imghdr
import random
Expand All @@ -15,6 +14,7 @@
import enum
import subprocess
import asyncio
import threading

from imgurpython import ImgurClient
import yaml
Expand Down Expand Up @@ -61,6 +61,16 @@
ERRORS_CNT_LIMIT = 2


HTTP_TIMEOUT = (5, 20)
HEAD_TIMEOUT = (3, 8)
TELEGRAM_TIMEOUT = 30
REDDIT_TIMEOUT = 20


_MONGO_DATABASES = {}
_MONGO_LOCK = threading.Lock()


@enum.unique
class SupplyResult(enum.Enum):
SUCCESSFULLY = 0
Expand All @@ -77,18 +87,34 @@ def _normalize_reddit_media_url(url):
return url.replace('auto=webp', 'auto=jpg')


def _http_request(method, url, *, timeout=HTTP_TIMEOUT, **kwargs):
kwargs.setdefault('timeout', timeout)
return requests.request(method, url, **kwargs)


def _get_shared_database(config):
key = (config['db']['host'], config['db']['name'])
with _MONGO_LOCK:
if key not in _MONGO_DATABASES:
client = pymongo.MongoClient(host=key[0])
_MONGO_DATABASES[key] = client[key[1]]
return _MONGO_DATABASES[key]


def get_url(submission, mp4_instead_gif=True):
'''
return TYPE, URL
E.x.: return 'img', 'http://example.com/pic.png'
'''

def what_is_inside(url):
header = requests.head(url).headers
if 'Content-Type' in header:
return header['Content-Type']
else:
try:
response = _http_request('head', url, timeout=HEAD_TIMEOUT, allow_redirects=True)
except RequestException:
return ''
if 'Content-Type' in response.headers:
return response.headers['Content-Type']
return ''

# If reddit native gallery
if hasattr(submission, 'gallery_data'):
Expand Down Expand Up @@ -248,7 +274,7 @@ def imgur_album_to_story(images):
elif 'gfycat.com' in urlparse(url).netloc:
rname = re.findall(r'gfycat.com\/(?:detail\/)?(\w*)', url)[0]
try:
r = requests.get(GFYCAT_GET + rname)
r = _http_request('get', GFYCAT_GET + rname)
if r.status_code != 200:
logging.info('Gfy fail prevented!')
return TYPE_OTHER, url
Expand All @@ -257,7 +283,7 @@ def imgur_album_to_story(images):
return TYPE_GIF, urls['mp4Url']
else:
return TYPE_GIF, urls['max5mbGif']
except KeyError:
except (KeyError, RequestException):
logging.info('Gfy fail prevented!')
return TYPE_OTHER, url
else:
Expand All @@ -267,20 +293,23 @@ def imgur_album_to_story(images):
def download_file(url, filename):
# http://stackoverflow.com/questions/16694907/how-to-download-large-file-in-python-with-requests-py
# NOTE the stream=True parameter
r = requests.get(url, stream=True)
if r.status_code >= 400:
try:
with _http_request('get', url, stream=True) as r:
if r.status_code >= 400:
return False
chunk_counter = 0
chunk_size = 1024
with open(filename, 'wb') as f:
for chunk in r.iter_content(chunk_size=chunk_size):
if chunk: # filter out keep-alive new chunks
f.write(chunk)
#f.flush() commented by recommendation from J.F.Sebastian
chunk_counter += 1
# It is not possible to send greater than 50 MB via Telegram
if chunk_counter > TELEGRAM_VIDEO_LIMIT / chunk_size:
return False
except RequestException:
return False
chunk_counter = 0
chunk_size = 1024
with open(filename, 'wb') as f:
for chunk in r.iter_content(chunk_size=chunk_size):
if chunk: # filter out keep-alive new chunks
f.write(chunk)
#f.flush() commented by recommendation from J.F.Sebastian
chunk_counter += 1
# It is not possible to send greater than 50 MB via Telegram
if chunk_counter > TELEGRAM_VIDEO_LIMIT / chunk_size:
return False
return True


Expand Down Expand Up @@ -314,20 +343,23 @@ def clean_after_module(submodule_name=None):

def md5_sum_from_url(url):
try:
r = requests.get(url, stream=True)
r = _http_request('get', url, stream=True)
except InvalidSchema:
return None
except MissingSchema:
return None
except RequestException:
return None
chunk_counter = 0
hash_store = hashlib.md5()
for chunk in r.iter_content(chunk_size=1024):
if chunk: # filter out keep-alive new chunks
hash_store.update(chunk)
chunk_counter += 1
# It is not possible to send greater than 50 MB via Telegram
if chunk_counter > 50 * 1024:
return None
with r:
for chunk in r.iter_content(chunk_size=1024):
if chunk: # filter out keep-alive new chunks
hash_store.update(chunk)
chunk_counter += 1
# It is not possible to send greater than 50 MB via Telegram
if chunk_counter > 50 * 1024:
return None
return hash_store.hexdigest()


Expand All @@ -343,12 +375,11 @@ def weighted_random_subreddit(weights):

def get_url_size(url):
# https://stackoverflow.com/questions/55226378/how-can-i-get-the-file-size-from-a-link-without-downloading-it-in-python
req = urllib.request.Request(url, method='HEAD')
try:
f = urllib.request.urlopen(req)
except Exception:
response = _http_request('head', url, timeout=HEAD_TIMEOUT, allow_redirects=True)
except RequestException:
return 0
content_length = f.headers.get('Content-Length')
content_length = response.headers.get('Content-Length')
if not content_length:
return 0
return int(content_length)
Expand All @@ -364,7 +395,13 @@ def __init__(self, t_channel=None, config=None):
with open(os.path.join('configs', 'prod.yml')) as f:
config = yaml.safe_load(f.read())
self.config = config
request = HTTPXRequest(connection_pool_size=8, pool_timeout=30)
request = HTTPXRequest(
connection_pool_size=8,
pool_timeout=TELEGRAM_TIMEOUT,
read_timeout=TELEGRAM_TIMEOUT,
write_timeout=TELEGRAM_TIMEOUT,
connect_timeout=TELEGRAM_TIMEOUT,
)
self.telegram_bot = Bot(self.config['telegram']['token'], request=request)
self._loop = asyncio.new_event_loop()
if t_channel is None:
Expand All @@ -385,12 +422,13 @@ def _run_async(self, coro):
return self._loop.run_until_complete(coro)

def _make_mongo_connections(self):
self.stats = pymongo.MongoClient(host=self.config['db']['host'])[self.config['db']['name']]['stats']
self.urls = pymongo.MongoClient(host=self.config['db']['host'])[self.config['db']['name']]['urls']
self.contents = pymongo.MongoClient(host=self.config['db']['host'])[self.config['db']['name']]['contents']
self.errors = pymongo.MongoClient(host=self.config['db']['host'])[self.config['db']['name']]['errors']
self.tasks = pymongo.MongoClient(host=self.config['db']['host'])[self.config['db']['name']]['tasks']
self.settings = pymongo.MongoClient(host=self.config['db']['host'])[self.config['db']['name']]['settings']
database = _get_shared_database(self.config)
self.stats = database['stats']
self.urls = database['urls']
self.contents = database['contents']
self.errors = database['errors']
self.tasks = database['tasks']
self.settings = database['settings']

def _get_file_name(self, ext='file'):
os.makedirs(TEMP_FOLDER, exist_ok=True)
Expand Down Expand Up @@ -536,7 +574,7 @@ def send_gif(self, url, text, parse_mode=None):

def _get_dash_audio_url(self, dash_url):
try:
resp = requests.get(dash_url)
resp = _http_request('get', dash_url)
if resp.status_code >= 400:
return None
base_urls = re.findall(r'<BaseURL>([^<]+)</BaseURL>', resp.text)
Expand Down
92 changes: 92 additions & 0 deletions tests/test_resilience.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,92 @@
import time
from types import SimpleNamespace
from unittest import mock

import requests

import sys
from pathlib import Path


REPO_ROOT = Path(__file__).resolve().parent.parent
APP_DIR = REPO_ROOT / 'reddit2telegram'
sys.path.insert(0, str(APP_DIR))

import task_queue
import utils


class FakeCollection:
def __init__(self, modified_count=0):
self.modified_count = modified_count
self.calls = []

def update_many(self, selector, update):
self.calls.append((selector, update))
return SimpleNamespace(modified_count=self.modified_count)


def test_recover_abandoned_tasks_requeues_stale_non_terminal_work():
collection = FakeCollection(modified_count=3)

with mock.patch('task_queue.time.time', return_value=1_000):
task_queue.recover_abandoned_tasks(collection)

assert len(collection.calls) == 1
selector, update = collection.calls[0]
assert selector['status']['$in'] == [
task_queue.TaskStatus.IN_PROGRESS.value,
task_queue.TaskStatus.SCHEDULED.value,
]
assert selector['updated_at']['$lt'] == 940
assert update['$set']['status'] == task_queue.TaskStatus.NEW.value
assert update['$set']['updated_at'] == 1_000


def test_get_url_returns_other_when_head_request_fails():
submission = SimpleNamespace(
url='https://example.com/image.jpg',
is_video=False,
media=None,
crosspost_parent_list=[],
is_self=False,
)

with mock.patch(
'utils._http_request',
side_effect=requests.RequestException('timeout'),
):
what, url = utils.get_url(submission)

assert what == utils.TYPE_OTHER
assert url == 'https://example.com/image.jpg'


def test_get_url_size_returns_zero_when_head_request_fails():
with mock.patch(
'utils._http_request',
side_effect=requests.RequestException('timeout'),
):
assert utils.get_url_size('https://example.com/file.mp4') == 0


def test_shared_database_reuses_single_mongo_client():
utils._MONGO_DATABASES.clear()
calls = []

class FakeClient:
def __init__(self, host):
calls.append(host)

def __getitem__(self, name):
return {'db_name': name}

config = {'db': {'host': 'localhost', 'name': 'reddit2telegram'}}

with mock.patch('utils.pymongo.MongoClient', FakeClient):
db_one = utils._get_shared_database(config)
db_two = utils._get_shared_database(config)

assert db_one == {'db_name': 'reddit2telegram'}
assert db_two == db_one
assert calls == ['localhost']