Middleware that handles JWT token refresh automatically.
| 38 | |
| 39 | |
| 40 | class TokenRefreshMiddleware: |
| 41 | """Middleware that handles JWT token refresh automatically.""" |
| 42 | |
| 43 | def __init__(self, token_endpoint: str, refresh_token: str) -> None: |
| 44 | self.token_endpoint = token_endpoint |
| 45 | self.refresh_token = refresh_token |
| 46 | self.access_token: str | None = None |
| 47 | self.token_expires_at: float | None = None |
| 48 | self._refresh_lock = asyncio.Lock() |
| 49 | |
| 50 | async def _refresh_access_token(self, session: ClientSession) -> str: |
| 51 | """Refresh the access token using the refresh token.""" |
| 52 | async with self._refresh_lock: |
| 53 | # Check if another coroutine already refreshed the token |
| 54 | if ( |
| 55 | self.token_expires_at |
| 56 | and time.time() < self.token_expires_at |
| 57 | and self.access_token |
| 58 | ): |
| 59 | _LOGGER.debug("Token already refreshed by another request") |
| 60 | return self.access_token |
| 61 | |
| 62 | _LOGGER.info("Refreshing access token...") |
| 63 | |
| 64 | # Make refresh request without middleware to avoid recursion |
| 65 | async with session.post( |
| 66 | self.token_endpoint, |
| 67 | json={"refresh_token": self.refresh_token}, |
| 68 | middlewares=(), # Disable middleware for this request |
| 69 | ) as resp: |
| 70 | resp.raise_for_status() |
| 71 | data = await resp.json() |
| 72 | |
| 73 | if "access_token" not in data: |
| 74 | raise ValueError("No access_token in refresh response") |
| 75 | |
| 76 | self.access_token = data["access_token"] |
| 77 | # Token expires in 5 minutes for demo, refresh 30 seconds early |
| 78 | expires_in = data.get("expires_in", 300) |
| 79 | self.token_expires_at = time.time() + expires_in - 30 |
| 80 | |
| 81 | _LOGGER.info( |
| 82 | "Token refreshed successfully, expires in %s seconds", expires_in |
| 83 | ) |
| 84 | if TYPE_CHECKING: |
| 85 | assert self.access_token is not None # Just assigned above |
| 86 | return self.access_token |
| 87 | |
| 88 | async def __call__( |
| 89 | self, |
| 90 | request: ClientRequest, |
| 91 | handler: ClientHandlerType, |
| 92 | ) -> ClientResponse: |
| 93 | """Add auth token to request, refreshing if needed.""" |
| 94 | # Skip token for refresh endpoint to avoid recursion |
| 95 | if str(request.url).endswith("/token/refresh"): |
| 96 | return await handler(request) |
| 97 |