From ef57d424252de47af05121bc1920ce09371532bc Mon Sep 17 00:00:00 2001 From: Florian Bezannier Date: Thu, 11 Nov 2021 23:32:54 +0100 Subject: [PATCH] handle rate limit --- charge_control.py | 6 +++++- libs/utils.py | 20 +++++++++++++------- my_psacc.py | 15 +++++++++------ test/test_unit.py | 14 +++++++++++++- 4 files changed, 40 insertions(+), 15 deletions(-) diff --git a/charge_control.py b/charge_control.py index 6c01c0f..049902a 100644 --- a/charge_control.py +++ b/charge_control.py @@ -8,6 +8,7 @@ from time import sleep import pytz from libs.psa_constants import DISCONNECTED, INPROGRESS, FINISHED +from libs.utils import RateLimitException from my_psacc import MyPSACC from mylogger import logger @@ -60,7 +61,10 @@ class ChargeControl: else: wakeup_timeout = self.wakeup_timeout if (datetime.utcnow().replace(tzinfo=pytz.UTC) - last_update).total_seconds() > 60 * wakeup_timeout: - self.psacc.wakeup(self.vin) + try: + self.psacc.wakeup(self.vin) + except RateLimitException: + logger.exception("force_update:") def process(self): now = datetime.now() diff --git a/libs/utils.py b/libs/utils.py index 2ad60c9..9db078e 100644 --- a/libs/utils.py +++ b/libs/utils.py @@ -26,19 +26,25 @@ def get_temp(latitude: str, longitude: str, api_key: str) -> float: return None +class RateLimitException(Exception): + pass + + def rate_limit(limit, every): def limit_decorator(func): semaphore = Semaphore(limit) @wraps(func) def wrapper(*args, **kwargs): - semaphore.acquire() - try: - return func(*args, **kwargs) - finally: # don't catch but ensure semaphore release - timer = Timer(every, semaphore.release) - timer.setDaemon(True) # allows the timer to be canceled on exit - timer.start() + if semaphore.acquire(blocking=False): + try: + return func(*args, **kwargs) + finally: # don't catch but ensure semaphore release + timer = Timer(every, semaphore.release) + timer.setDaemon(True) # allows the timer to be canceled on exit + timer.start() + else: + raise RateLimitException return wrapper diff --git a/my_psacc.py b/my_psacc.py index 67db66c..ef03c12 100644 --- a/my_psacc.py +++ b/my_psacc.py @@ -23,7 +23,7 @@ from otp.otp import load_otp, new_otp_session, save_otp, ConfigException, Otp from psa_connectedcar.rest import ApiException from mylogger import logger -from libs.utils import rate_limit, parse_hour +from libs.utils import rate_limit, parse_hour, RateLimitException from web.abrp import Abrp from web.db import Database @@ -230,8 +230,8 @@ class MyPSACC: last_update: datetime = self.remote_token_last_update if (datetime.now() - last_update).total_seconds() < MQTT_TOKEN_TTL: return res - self.refresh_token() try: + self.refresh_token() if bad_remote_token: logger.error("remote_refresh_token isn't defined") else: @@ -255,8 +255,8 @@ class MyPSACC: self.mqtt_client.username_pw_set("IMA_OAUTH_ACCESS_TOKEN", self.remote_access_token) self.save_config() return res - except RequestException as e: - logger.error("Can't refresh remote token %s", e) + except (RequestException, RateLimitException) as e: + logger.exception("Can't refresh remote token %s", e) sleep(60) return None @@ -307,7 +307,7 @@ class MyPSACC: logger.warning("charge begin but API isn't updated") sleep(60) self.wakeup(data["vin"]) - except (IndexError, AttributeError): + except (IndexError, AttributeError, RateLimitException): logger.exception("on_mqtt_message:") except KeyError: logger.exception("on_mqtt_message:") @@ -330,7 +330,10 @@ class MyPSACC: def __keep_mqtt(self): # avoid token expiration timeout = 3600 * 24 # 1 day if len(self.vehicles_list) > 0: - self.wakeup(self.vehicles_list[0].vin) + try: + self.wakeup(self.vehicles_list[0].vin) + except RateLimitException: + logger.exception("__keep_mqtt") t = threading.Timer(timeout, self.__keep_mqtt) t.setDaemon(True) t.start() diff --git a/test/test_unit.py b/test/test_unit.py index c7f6e2b..1a0454f 100644 --- a/test/test_unit.py +++ b/test/test_unit.py @@ -22,7 +22,7 @@ from charge_control import ChargeControls from test.utils import DATA_DIR, record_position, latitude, longitude, date0, date1, date2, date3, record_charging, \ vehicule_list, get_new_test_db, get_date, date4 from trip import Trips -from libs.utils import get_temp, parse_hour +from libs.utils import get_temp, parse_hour, rate_limit, RateLimitException from web.db import Database from web.figures import get_figures, get_battery_curve_fig, get_altitude_fig from deepdiff import DeepDiff @@ -253,6 +253,18 @@ class TestUnit(unittest.TestCase): expected_res = [[2, 0, 0], [3, 14, 0], [0, 0, 2], [0, 30, 0]] assert expected_res == [parse_hour(h) for h in ["PT2H", "PT3H14", "PT2S", "PT30M"]] + def test_rate_limit(self): + @rate_limit(2, 10) + def test_fct(): + pass + test_fct() + test_fct() + try: + test_fct() + raise Exception("It should have raise RateLimitException") + except RateLimitException: + pass + if __name__ == '__main__': my_logger(handler_level=os.environ.get("DEBUG_LEVEL", 20))