import time from dataclasses import dataclass, field import requests from smelog.factory import BoundLogger from src.config import ATLAS_API_HOST, DNA_API_HOST, DNA_ATLAS_AUDIENCE, DNA_ATLAS_CLIENT_ID, DNA_ATLAS_SECRET @dataclass class Response: status_code: int = 0 content: dict = field(default_factory=dict) class API_Client: TOKEN = None EXPIRES_AT = None API_HOST = None ATLAS_CLIENT_ID = None ATLAS_SECRET = None ATLAS_AUDIENCE = None def __init__(self, logger: BoundLogger): self.logger = logger def get_token(self): if self.TOKEN and self.EXPIRES_AT > time.time(): return self.TOKEN payload = { "client_id": self.ATLAS_CLIENT_ID, "client_secret": self.ATLAS_SECRET, "grant_type": "client_credentials", "audience": self.ATLAS_AUDIENCE, } i = 1 r = None while (not hasattr(r, "status_code") or r.status_code != 200) and i <= 5: time.sleep(i - 1) self.logger.info(f"Atlas token request sent #{i}") r = requests.post(ATLAS_API_HOST + "oauth/token", json=payload) i += 1 r = r.json() self.TOKEN = r["access_token"] self.EXPIRES_AT = r["expires_in"] + time.time() return self.TOKEN def request(self, endpoint: str, payload: dict = None): headers = {"Authorization": "Bearer " + self.get_token()} r = requests.get(self.API_HOST + endpoint, params=payload or {}, headers=headers) kwargs = {"status_code": r.status_code} if r.status_code == 200: kwargs["content"] = r.json() return Response(**kwargs) class DNA_API_Client(API_Client): API_HOST = DNA_API_HOST ATLAS_CLIENT_ID = DNA_ATLAS_CLIENT_ID ATLAS_SECRET = DNA_ATLAS_SECRET ATLAS_AUDIENCE = DNA_ATLAS_AUDIENCE