diff --git a/telebot/__init__.py b/telebot/__init__.py index 45c1da00a..c466ad425 100644 --- a/telebot/__init__.py +++ b/telebot/__init__.py @@ -1171,15 +1171,15 @@ def infinity_polling(self, timeout: Optional[int]=20, skip_pending: Optional[boo logger_level=logger_level, allowed_updates=allowed_updates, restart_on_change=False, *args, **kwargs) except Exception as e: - if logger_level and logger_level >= logging.ERROR: + if logger_level and logger_level <= logging.ERROR: logger.error("Infinity polling exception: %s", self.__hide_token(str(e))) - if logger_level and logger_level >= logging.DEBUG: + if logger_level and logger_level <= logging.DEBUG: logger.error("Exception traceback:\n%s", self.__hide_token(traceback.format_exc())) time.sleep(3) continue - if logger_level and logger_level >= logging.INFO: + if logger_level and logger_level <= logging.INFO: logger.error("Infinity polling: polling exited") - if logger_level and logger_level >= logging.INFO: + if logger_level and logger_level <= logging.INFO: logger.error("Break infinity polling") @@ -1274,7 +1274,7 @@ def __threaded_polling(self, non_stop = False, interval = 0, timeout = None, lon warning = "\n Warning: this message appearance will be changed. Set logger_level=logging.INFO to continue seeing it." else: warning = "" - #if logger_level and logger_level >= logging.INFO: # enable in future releases. Change output to logger.error + #if logger_level and logger_level <= logging.INFO: # enable in future releases. Change output to logger.error logger.info('Started polling.' + warning) self.__stop_polling.clear() error_interval = 0.25 @@ -1297,16 +1297,16 @@ def __threaded_polling(self, non_stop = False, interval = 0, timeout = None, lon except apihelper.ApiException as e: handled = self._handle_exception(e) if not handled: - if logger_level and logger_level >= logging.ERROR: + if logger_level and logger_level <= logging.ERROR: logger.error("Threaded polling exception: %s", self.__hide_token(str(e))) - if logger_level and logger_level >= logging.DEBUG: + if logger_level and logger_level <= logging.DEBUG: logger.error("Exception traceback:\n%s", self.__hide_token(traceback.format_exc())) if not non_stop: self.__stop_polling.set() - # if logger_level and logger_level >= logging.INFO: # enable in future releases. Change output to logger.error + # if logger_level and logger_level <= logging.INFO: # enable in future releases. Change output to logger.error logger.info("Exception occurred. Stopping." + warning) else: - # if logger_level and logger_level >= logging.INFO: # enable in future releases. Change output to logger.error + # if logger_level and logger_level <= logging.INFO: # enable in future releases. Change output to logger.error logger.info("Waiting for {0} seconds until retry".format(error_interval) + warning) time.sleep(error_interval) if error_interval * 2 < 60: @@ -1320,7 +1320,7 @@ def __threaded_polling(self, non_stop = False, interval = 0, timeout = None, lon polling_thread.clear_exceptions() #* self.worker_pool.clear_exceptions() #* except KeyboardInterrupt: - # if logger_level and logger_level >= logging.INFO: # enable in future releases. Change output to logger.error + # if logger_level and logger_level <= logging.INFO: # enable in future releases. Change output to logger.error logger.info("KeyboardInterrupt received." + warning) self.__stop_polling.set() break @@ -1339,7 +1339,7 @@ def __threaded_polling(self, non_stop = False, interval = 0, timeout = None, lon polling_thread.stop() polling_thread.clear_exceptions() self.worker_pool.clear_exceptions() - #if logger_level and logger_level >= logging.INFO: # enable in future releases. Change output to logger.error + #if logger_level and logger_level <= logging.INFO: # enable in future releases. Change output to logger.error logger.info('Stopped polling.' + warning) @@ -1349,7 +1349,7 @@ def __non_threaded_polling(self, non_stop=False, interval=0, timeout=None, long_ warning = "\n Warning: this message appearance will be changed. Set logger_level=logging.INFO to continue seeing it." else: warning = "" - #if logger_level and logger_level >= logging.INFO: # enable in future releases. Change output to logger.error + #if logger_level and logger_level <= logging.INFO: # enable in future releases. Change output to logger.error logger.info('Started polling.' + warning) self.__stop_polling.clear() error_interval = 0.25 @@ -1361,23 +1361,23 @@ def __non_threaded_polling(self, non_stop=False, interval=0, timeout=None, long_ except apihelper.ApiException as e: handled = self._handle_exception(e) if not handled: - if logger_level and logger_level >= logging.ERROR: + if logger_level and logger_level <= logging.ERROR: logger.error("Polling exception: %s", str(e)) - if logger_level and logger_level >= logging.DEBUG: + if logger_level and logger_level <= logging.DEBUG: logger.error("Exception traceback:\n%s", traceback.format_exc()) if not non_stop: self.__stop_polling.set() - # if logger_level and logger_level >= logging.INFO: # enable in future releases. Change output to logger.error + # if logger_level and logger_level <= logging.INFO: # enable in future releases. Change output to logger.error logger.info("Exception occurred. Stopping." + warning) else: - # if logger_level and logger_level >= logging.INFO: # enable in future releases. Change output to logger.error + # if logger_level and logger_level <= logging.INFO: # enable in future releases. Change output to logger.error logger.info("Waiting for {0} seconds until retry".format(error_interval) + warning) time.sleep(error_interval) error_interval *= 2 else: time.sleep(error_interval) except KeyboardInterrupt: - # if logger_level and logger_level >= logging.INFO: # enable in future releases. Change output to logger.error + # if logger_level and logger_level <= logging.INFO: # enable in future releases. Change output to logger.error logger.info("KeyboardInterrupt received." + warning) self.__stop_polling.set() break @@ -1387,7 +1387,7 @@ def __non_threaded_polling(self, non_stop=False, interval=0, timeout=None, long_ raise e else: time.sleep(error_interval) - #if logger_level and logger_level >= logging.INFO: # enable in future releases. Change output to logger.error + #if logger_level and logger_level <= logging.INFO: # enable in future releases. Change output to logger.error logger.info('Stopped polling.' + warning) diff --git a/telebot/apihelper.py b/telebot/apihelper.py index 25180e562..a687066f1 100644 --- a/telebot/apihelper.py +++ b/telebot/apihelper.py @@ -12,12 +12,26 @@ from requests.exceptions import HTTPError, ConnectionError, Timeout from requests.adapters import HTTPAdapter + +def _get_multipart_header_formatter(urllib3_fields): + """Return urllib3's active multipart-header formatter and its name.""" + try: + return ( + urllib3_fields.format_multipart_header_param, + 'format_multipart_header_param', + ) + except AttributeError: + return urllib3_fields.format_header_param, 'format_header_param' + + try: # noinspection PyUnresolvedReferences from requests.packages.urllib3 import fields - format_header_param = fields.format_header_param -except ImportError: + format_header_param, format_header_param_name = _get_multipart_header_formatter(fields) +except (ImportError, AttributeError): + fields = None format_header_param = None + format_header_param_name = None import telebot from telebot import types from telebot import util @@ -100,7 +114,7 @@ def _make_request(token, method_name, method='get', params=None, files=None): if files and format_header_param: - fields.format_header_param = _no_encode(format_header_param) + setattr(fields, format_header_param_name, _no_encode(format_header_param)) if params: if 'timeout' in params: read_timeout = params.pop('timeout') diff --git a/telebot/async_telebot.py b/telebot/async_telebot.py index 453ce5e60..f3d2a85ef 100644 --- a/telebot/async_telebot.py +++ b/telebot/async_telebot.py @@ -386,15 +386,15 @@ async def infinity_polling(self, timeout: Optional[int]=20, skip_pending: Option await self._process_polling(non_stop=True, timeout=timeout, request_timeout=request_timeout, allowed_updates=allowed_updates, *args, **kwargs) except Exception as e: - if logger_level and logger_level >= logging.ERROR: + if logger_level and logger_level <= logging.ERROR: logger.error("Infinity polling exception: %s", self.__hide_token(str(e))) - if logger_level and logger_level >= logging.DEBUG: + if logger_level and logger_level <= logging.DEBUG: logger.error("Exception traceback:\n%s", self.__hide_token(traceback.format_exc())) await asyncio.sleep(3) continue - if logger_level and logger_level >= logging.INFO: + if logger_level and logger_level <= logging.INFO: logger.error("Infinity polling: polling exited") - if logger_level and logger_level >= logging.INFO: + if logger_level and logger_level <= logging.INFO: logger.error("Break infinity polling") async def _handle_exception(self, exception: Exception) -> bool: diff --git a/tests/test_apihelper_95.py b/tests/test_apihelper_95.py index bf0f6c304..80e1d26f4 100644 --- a/tests/test_apihelper_95.py +++ b/tests/test_apihelper_95.py @@ -1,6 +1,57 @@ +from types import SimpleNamespace + from telebot import apihelper +def test_get_multipart_header_formatter_prefers_current_urllib3_name(): + def current_formatter(key, value): + return 'current={0}'.format(value) + + def legacy_formatter(key, value): + return 'legacy={0}'.format(value) + + fields = SimpleNamespace( + format_multipart_header_param=current_formatter, + format_header_param=legacy_formatter, + ) + + formatter, name = apihelper._get_multipart_header_formatter(fields) + + assert formatter is current_formatter + assert name == 'format_multipart_header_param' + + +def test_get_multipart_header_formatter_supports_legacy_urllib3_name(): + def legacy_formatter(key, value): + return 'legacy={0}'.format(value) + + fields = SimpleNamespace(format_header_param=legacy_formatter) + + formatter, name = apihelper._get_multipart_header_formatter(fields) + + assert formatter is legacy_formatter + assert name == 'format_header_param' + + +def test_make_request_patches_selected_multipart_header_formatter(monkeypatch): + def formatter(key, value): + return '{0}="{1}"'.format(key, value) + + fields = SimpleNamespace(format_multipart_header_param=formatter) + monkeypatch.setattr(apihelper, 'fields', fields) + monkeypatch.setattr(apihelper, 'format_header_param', formatter) + monkeypatch.setattr(apihelper, 'format_header_param_name', 'format_multipart_header_param') + response = SimpleNamespace( + text='{"ok": true, "result": true}', + status_code=200, + json=lambda: {'ok': True, 'result': True}, + ) + monkeypatch.setattr(apihelper, 'CUSTOM_REQUEST_SENDER', lambda *args, **kwargs: response) + + assert apihelper._make_request('token', 'test', method='post', files={'document': ('test.txt', object())}) is True + assert fields.format_multipart_header_param('filename', 'test file.txt') == 'filename=test file.txt' + + def test_promote_chat_member_can_manage_tags(monkeypatch): captured = {} diff --git a/tests/test_async_telebot.py b/tests/test_async_telebot.py index ceaa33fbe..2b72da6c0 100644 --- a/tests/test_async_telebot.py +++ b/tests/test_async_telebot.py @@ -5,8 +5,12 @@ network I/O. """ import asyncio +import logging + +import pytest from telebot import types +import telebot.async_telebot as async_telebot from telebot.async_telebot import AsyncTeleBot @@ -19,6 +23,46 @@ def _make_fake_me() -> types.User: }) +@pytest.mark.parametrize( + 'logger_level, expected_count, includes_traceback', + [ + (logging.DEBUG, 3, True), + (logging.INFO, 2, False), + (logging.ERROR, 1, False), + (None, 0, False), + ], +) +def test_infinity_polling_honors_logger_level( + monkeypatch, logger_level, expected_count, includes_traceback): + class RecordingLogger: + def __init__(self): + self.messages = [] + + def error(self, message, *args): + self.messages.append(message % args if args else message) + + bot = AsyncTeleBot('1:fake', validate_token=False) + logger = RecordingLogger() + + async def fail_polling(*args, **kwargs): + bot._polling = False + raise RuntimeError('polling failed') + + async def no_sleep(_seconds): + return None + + monkeypatch.setattr(async_telebot, 'logger', logger) + monkeypatch.setattr(async_telebot.asyncio, 'sleep', no_sleep) + monkeypatch.setattr(bot, '_process_polling', fail_polling) + + asyncio.run(bot.infinity_polling(logger_level=logger_level)) + + assert len(logger.messages) == expected_count + if logger_level: + assert logger.messages[0] == 'Infinity polling exception: polling failed' + assert any('Exception traceback:' in message for message in logger.messages) is includes_traceback + + def test_process_polling_retains_update_processing_tasks(): """Regression test for issue #2572. diff --git a/tests/test_telebot.py b/tests/test_telebot.py index 14d80e89b..0a3cf0e0b 100644 --- a/tests/test_telebot.py +++ b/tests/test_telebot.py @@ -1,6 +1,7 @@ # -*- coding: utf-8 -*- import sys import warnings +import logging sys.path.append('../') @@ -34,6 +35,44 @@ def deprecated2_new_function(): def deprecated2_old_function(): print("deprecated2_old_function") + +@pytest.mark.parametrize( + 'logger_level, expected_count, includes_traceback', + [ + (logging.DEBUG, 3, True), + (logging.INFO, 2, False), + (logging.ERROR, 1, False), + (None, 0, False), + ], +) +def test_infinity_polling_honors_logger_level( + monkeypatch, logger_level, expected_count, includes_traceback): + class RecordingLogger: + def __init__(self): + self.messages = [] + + def error(self, message, *args): + self.messages.append(message % args if args else message) + + bot = telebot.TeleBot('1:fake', validate_token=False) + logger = RecordingLogger() + + def fail_polling(*args, **kwargs): + bot._TeleBot__stop_polling.set() + raise RuntimeError('polling failed') + + monkeypatch.setattr(telebot, 'logger', logger) + monkeypatch.setattr(telebot.time, 'sleep', lambda _seconds: None) + monkeypatch.setattr(bot, 'polling', fail_polling) + + bot.infinity_polling(logger_level=logger_level) + + assert len(logger.messages) == expected_count + if logger_level: + assert logger.messages[0] == 'Infinity polling exception: polling failed' + assert any('Exception traceback:' in message for message in logger.messages) is includes_traceback + + @pytest.mark.skipif(should_skip, reason="No environment variables configured") class TestTeleBot: def test_message_listener(self):