from __future__ import annotations import abc import logging import os from typing import Any, Dict, Optional, cast from urllib.parse import urljoin import cachelib import httpx from authlib.jose import JWTClaims, jwt from authlib.oauth2.rfc6749 import MissingAuthorizationError, UnsupportedTokenTypeError from authlib.oauth2.rfc6750 import InvalidTokenError from owslib.enums import Environment QA_ISSUERS = ["https://qalogin.theorchard.com/", "https://qa-orchard.auth0.com/"] PROD_ISSUERS = ["https://login.distroauth.com/", "https://workstation.auth0.com/"] QA_AUDIENCE = "https://qa-ows.theorchard.io" PROD_AUDIENCE = "https://prod-ows.theorchard.io" QA_API_AUDIENCE = "https://workstation.qaorch.com/api" PROD_API_AUDIENCE = "https://workstation.theorchard.com/api" DEFAULT_CACHE_TIMEOUT = 60 * 60 # 1 hour logger = logging.getLogger(__name__) class BaseJWTAuth(abc.ABC): TOKEN_TYPE = "bearer" _key = "_jwks_set" def __init__( self, jwks_url: str, claims_options: Optional[Dict[str, Any]] = None, cache_timeout: Optional[int] = DEFAULT_CACHE_TIMEOUT, realm: Optional[str] = None, cache: Optional[cachelib.SimpleCache] = None, ): self.jwks_url = jwks_url self.claims_options = claims_options or {} self.cache_timeout = cache_timeout self.realm = realm self.cache = cache or cachelib.SimpleCache() def get_token_string_from_authorization(self, authorization: Optional[str]) -> str: if not authorization: raise MissingAuthorizationError token_parts = authorization.split(maxsplit=1) if len(token_parts) != 2: raise UnsupportedTokenTypeError token_type, token_string = token_parts if token_type.lower() != self.TOKEN_TYPE: raise UnsupportedTokenTypeError return token_string def get_token(self, token_string: str, key: Any) -> JWTClaims: try: return self.decode_token(token_string, key=key) except Exception as exc: logger.warning(f"Failed to decode JWT: {exc}", exc_info=True) raise InvalidTokenError(realm=self.realm) from exc def decode_token(self, token_string: str, key: Any) -> JWTClaims: token = jwt.decode(token_string, key=key, claims_options=self.claims_options) token.validate() return token class JWTAuth(BaseJWTAuth): def authenticate(self, authorization: Optional[str]) -> JWTClaims: token_string = self.get_token_string_from_authorization(authorization) return self.get_token(token_string, key=self.get_jwk_set()) def _get_jwk_set_from_cache(self) -> Optional[Dict[str, Any]]: logger.info("Fetching JWK set from cache ...") if token := self.cache.get(self._key): return cast(Dict[str, Any], token) return None def _get_jwk_set(self) -> Dict[str, Any]: logger.info(f"Fetching JWK set from {self.jwks_url} ...") response = httpx.get(self.jwks_url) response.raise_for_status() jwk_set = response.json() if not isinstance(jwk_set, dict): raise ValueError( f"Invalid jwk_set type. Expected dict, got {jwk_set} instead." ) return jwk_set def get_jwk_set(self) -> Dict[str, Any]: if not (jwk_set := self._get_jwk_set_from_cache()): jwk_set = self._get_jwk_set() self.cache.set(self._key, value=jwk_set, timeout=self.cache_timeout) return jwk_set class AsyncJWTAuth(BaseJWTAuth): async def authenticate(self, authorization: Optional[str]) -> JWTClaims: token_string = self.get_token_string_from_authorization(authorization) return self.get_token(token_string, key=await self.get_jwk_set()) async def _get_jwk_set_from_cache(self) -> Optional[Dict[str, Any]]: logger.info("Fetching JWK set from cache ...") if token := self.cache.get(self._key): return cast(Dict[str, Any], token) return None async def _get_jwk_set(self) -> Dict[str, Any]: logger.info(f"Fetching jwks from {self.jwks_url} ...") async with httpx.AsyncClient() as client: response = await client.get(self.jwks_url) response.raise_for_status() jwk_set = response.json() if not isinstance(jwk_set, dict): raise ValueError( f"Invalid jwk_set type. Expected dict, got {jwk_set} instead." ) return jwk_set async def get_jwk_set(self) -> Dict[str, Any]: if not (jwk_set := await self._get_jwk_set_from_cache()): jwk_set = await self._get_jwk_set() self.cache.set(self._key, value=jwk_set, timeout=self.cache_timeout) return jwk_set def get_default_jws_url(environment: Environment) -> str: issuers_env = os.environ.get("AUTH_ISSUERS") if issuers_env: issuers = issuers_env.replace(" ", "").split(",") else: issuers = PROD_ISSUERS if environment == Environment.PROD else QA_ISSUERS return urljoin(issuers[0], ".well-known/jwks.json") def get_default_claims_options(environment: Environment) -> Dict[str, Any]: m2m_api_audience = PROD_AUDIENCE if environment == Environment.PROD else QA_AUDIENCE orchard_api_audience = ( PROD_API_AUDIENCE if environment == Environment.PROD else QA_API_AUDIENCE ) return { "aud": { "essential": True, "values": [m2m_api_audience, orchard_api_audience], }, "iss": {"essential": True}, "sub": {"essential": True}, "exp": {"essential": True}, "iat": {"essential": True}, "azp": {"essential": True}, } def jwt_auth_from_config( environment: Environment, cache_timeout: Optional[int] = DEFAULT_CACHE_TIMEOUT, realm: Optional[str] = None, ) -> JWTAuth: return JWTAuth( jwks_url=get_default_jws_url(environment), claims_options=get_default_claims_options(environment), cache_timeout=cache_timeout, realm=realm, ) def async_jwt_auth_from_config( environment: Environment, cache_timeout: Optional[int] = DEFAULT_CACHE_TIMEOUT, realm: Optional[str] = None, ) -> AsyncJWTAuth: return AsyncJWTAuth( jwks_url=get_default_jws_url(environment), claims_options=get_default_claims_options(environment), cache_timeout=cache_timeout, realm=realm, ) def jwt_auth_enabled_for_env( environment: Environment, *, enabled: Optional[bool] ) -> bool: if enabled is None: return environment in [Environment.QA, Environment.PROD] return enabled