handle rate limit

This commit is contained in:
Florian Bezannier
2021-11-11 23:32:54 +01:00
parent 0de9cfbb00
commit ef57d42425
4 changed files with 40 additions and 15 deletions
+5 -1
View File
@@ -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()
+13 -7
View File
@@ -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
+9 -6
View File
@@ -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()
+13 -1
View File
@@ -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))