import config import pickle import redis from dataclasses import dataclass from typing import Dict, Optional from marshmallow import EXCLUDE from utils.decorators import redis_exception_handler from external_api.base.clients.auth_strategies.Base import AuthStrategy from external_api.base.clients.client import ApiClient @dataclass class AtlasTokenHolder: access_token: str expires_in: int class Meta: unknown = EXCLUDE @dataclass class AtlasCredentials: clientId: str clientSecret: str audience: str grantType: str = "client_credentials" class Meta: unknown = EXCLUDE class AtlasAuthStrategy(AuthStrategy): redis = redis.Redis(host=config.REDIS_HOST) token_holder: Optional[AtlasTokenHolder] = None http_client: ApiClient credentials: AtlasCredentials redis_token_key: str token_cache_lifespan: int def __init__( self, client: ApiClient, credentials: AtlasCredentials, redis_token_key: str, token_cache_lifespan: int, ): self.http_client = client self.credentials = credentials self.redis_token_key = redis_token_key self.token_cache_lifespan = token_cache_lifespan async def authorize(self, headers: Dict[str, str]) -> Dict[str, str]: token = await self.get_token() headers[self.AUTH_HEADER] = self.build_authorization_header(token) return headers async def get_token(self) -> str: self.token_holder = self.get_cached_token() if not self.token_holder: self.token_holder = await self.refresh_token() return self.token_holder.access_token @redis_exception_handler async def refresh_token(self) -> AtlasTokenHolder: token = await self.load_atlas_token() self.redis.set(self.redis_token_key, pickle.dumps(token), self.token_cache_lifespan) return token @redis_exception_handler def get_cached_token(self) -> Optional[AtlasTokenHolder]: token = None cached_token = self.redis.get(self.redis_token_key) if cached_token is not None: token = pickle.loads(cached_token) return token async def load_atlas_token(self) -> AtlasTokenHolder: params = { "client_id": self.credentials.clientId, "client_secret": self.credentials.clientSecret, "audience": self.credentials.audience, "grant_type": self.credentials.grantType, } return await self.http_client.post("/oauth/token", payload=params, response_type=AtlasTokenHolder)