from __future__ import annotations from jose import jwt, ExpiredSignatureError, JWTError from src.application.contracts import ILogger, IJwtService from src.application.domain.dto import AccessTokenPayload from src.application.domain.exceptions import ApplicationException from src.infrastructure.config.settings import settings from src.infrastructure.vault import JwtKeyStore class JwtService(IJwtService): def __init__(self, logger: ILogger, key_store: JwtKeyStore) -> None: self._logger = logger self._key_store = key_store async def decode_access_token(self, token: str) -> AccessTokenPayload: payload = await self._decode_and_verify(token) if payload.get('type') != 'access': self._logger.warning(f'Access token invalid type received_type={payload.get('type')}') raise ApplicationException(status_code=401, message='Invalid token type') try: return AccessTokenPayload( sub=str(payload['sub']), type='access', sid=str(payload['sid']), iat=int(payload['iat']), nbf=int(payload['nbf']), exp=int(payload['exp']), iss=payload.get('iss'), aud=payload.get('aud'), ) except KeyError as exception: self._logger.warning(f'Access token missing claim error={str(exception)}') raise ApplicationException(status_code=401, message=f'Missing token claim: {exception}') async def _decode_and_verify(self, token: str) -> dict: kid: str | None = None try: header = jwt.get_unverified_header(token) kid = header.get('kid') if not kid: self._logger.warning(f'JWT header missing kid header={header}') raise ApplicationException(status_code=401, message='Missing token header: kid') received_alg = header.get('alg') if received_alg != settings.JWT_ALGORITHM: self._logger.warning(f'JWT invalid algorithm kid={kid} received_alg={received_alg} expected_alg={settings.JWT_ALGORITHM}') raise ApplicationException(status_code=401, message='Invalid token algorithm') public_pem = await self._key_store.get_public_key_for_kid(str(kid)) if not public_pem: self._logger.info(f'JWT kid miss kid={kid} forcing keystore refresh') await self._key_store.refresh() public_pem = await self._key_store.get_public_key_for_kid(str(kid)) if not public_pem: self._logger.warning(f'JWT unknown kid kid={kid}') raise ApplicationException(status_code=401, message='Unknown token kid') options = { 'verify_signature': True, 'verify_exp': True, 'verify_nbf': True, 'verify_iat': True, 'verify_aud': bool(settings.JWT_AUDIENCE), 'verify_iss': bool(settings.JWT_ISSUER), 'require_exp': True, 'require_iat': True, 'require_nbf': True, 'require_sub': True, 'leeway': 10, } payload = jwt.decode( token, public_pem, algorithms=[settings.JWT_ALGORITHM], audience=settings.JWT_AUDIENCE or None, issuer=settings.JWT_ISSUER or None, options=options, ) if 'sid' not in payload: self._logger.warning(f'JWT missing sid claim kid={kid}') raise ApplicationException(status_code=401, message='Missing token claim: sid') if 'type' not in payload: self._logger.warning(f'JWT missing type claim kid={kid}') raise ApplicationException(status_code=401, message='Missing token claim: type') return payload except ExpiredSignatureError as exception: self._logger.info(f'JWT expired kid={kid} error={str(exception)}') raise ApplicationException(status_code=401, message='Token expired') except ApplicationException: raise except JWTError as exception: self._logger.warning(f'JWT decode failed kid={kid} error={str(exception)}') raise ApplicationException(status_code=401, message='Invalid token') except Exception as exception: self._logger.error(f'Unexpected JWT decode error kid={kid} error={str(exception)}') raise ApplicationException(status_code=500, message='JWT decode failed')