import datetime import logging from dataclasses import dataclass from typing import Any, Self, cast import dateutil.parser from anyio import to_thread from fansifter_common.protocols import AsyncCache, Cache, SecretsManager from fansifter_common.utils import timezone DEFAULT_EXPIRE_WINDOW = 60 logger = logging.getLogger(__name__) @dataclass class M2MToken: token_string: str expires_at: datetime.datetime @classmethod def create(cls, token_string: str, expires_at: str) -> Self: return cls( token_string=token_string, expires_at=dateutil.parser.isoparse(expires_at).replace( tzinfo=timezone.UTC ), ) def has_expired(self, window: int | None = None) -> bool: return self.expires_at < ( timezone.now() + datetime.timedelta(seconds=window or 0) ) class BaseM2MTokenManager: cache_key = "_m2m_token" def __init__( self, secrets_manager: SecretsManager, secret_name_key: str, secret_expire_name_key: str, expire_window: int = DEFAULT_EXPIRE_WINDOW, ) -> None: self.secret_name_key = secret_name_key self.secret_expire_name_key = secret_expire_name_key self.expire_window = expire_window self._secrets_manager = secrets_manager class M2MTokenManager(BaseM2MTokenManager): def __init__( self, secrets_manager: SecretsManager, secret_name_key: str, secret_expire_name_key: str, expire_window: int = DEFAULT_EXPIRE_WINDOW, cache: Cache | dict[str, Any] | None = None, ) -> None: super().__init__( secrets_manager=secrets_manager, secret_name_key=secret_name_key, secret_expire_name_key=secret_expire_name_key, expire_window=expire_window, ) self._cache = cache or {} def _get_token_payload_from_secret_manager(self) -> M2MToken: logger.info("Fetching M2M token from secrets manager ...") token_string = self._secrets_manager.get_secret(self.secret_name_key) expires_at_string = self._secrets_manager.get_secret( self.secret_expire_name_key ) return M2MToken.create(token_string=token_string, expires_at=expires_at_string) def _get_token_from_cache(self) -> M2MToken | None: logger.info("Fetching M2M token from cache ...") token = self._cache.get(self.cache_key) return cast(M2MToken | None, token) def get_token_string(self) -> str: token = self._get_token_from_cache() if (not token) or token and token.has_expired(window=self.expire_window): token = self._get_token_payload_from_secret_manager() # cache forever and rely always on token.expires_at if isinstance(self._cache, Cache): self._cache.set(self.cache_key, value=token, timeout=0) else: self._cache[self.cache_key] = token return token.token_string class AsyncM2MTokenManager(BaseM2MTokenManager): def __init__( self, secrets_manager: SecretsManager, secret_name_key: str, secret_expire_name_key: str, expire_window: int = DEFAULT_EXPIRE_WINDOW, cache: AsyncCache | dict[str, Any] | None = None, ) -> None: super().__init__( secrets_manager=secrets_manager, secret_name_key=secret_name_key, secret_expire_name_key=secret_expire_name_key, expire_window=expire_window, ) self._cache = cache or {} async def _get_token_payload_from_secret_manager(self) -> M2MToken: logger.info("Fetching M2M token from secrets manager ...") token_string = await to_thread.run_sync( self._secrets_manager.get_secret, self.secret_name_key ) expires_at_string = await to_thread.run_sync( self._secrets_manager.get_secret, self.secret_expire_name_key ) return M2MToken.create(token_string=token_string, expires_at=expires_at_string) async def _get_token_from_cache(self) -> M2MToken | None: logger.info("Fetching M2M token from cache ...") if isinstance(self._cache, AsyncCache): token = await self._cache.get(self.cache_key) else: token = self._cache.get(self.cache_key) return cast(M2MToken | None, token) async def get_token_string(self) -> str: token = await self._get_token_from_cache() if (not token) or token and token.has_expired(window=self.expire_window): token = await self._get_token_payload_from_secret_manager() if isinstance(self._cache, AsyncCache): await self._cache.set(self.cache_key, value=token, ttl=0) else: self._cache[self.cache_key] = token return token.token_string