diff --git a/mddocs/docs/changelog/next_release/410.improvement.md b/mddocs/docs/changelog/next_release/410.improvement.md new file mode 100644 index 00000000..24890ff3 --- /dev/null +++ b/mddocs/docs/changelog/next_release/410.improvement.md @@ -0,0 +1 @@ +Cache Keycloak certs for JWT token validation. Previously every API request with KeycloakAuthProvider enabled made a request to Keycloak. diff --git a/syncmaster/db/factory.py b/syncmaster/db/factory.py index 2aa33527..e0bf2542 100644 --- a/syncmaster/db/factory.py +++ b/syncmaster/db/factory.py @@ -1,23 +1,13 @@ # SPDX-FileCopyrightText: 2023-present MTS PJSC # SPDX-License-Identifier: Apache-2.0 -from collections.abc import AsyncGenerator, Callable -from typing import Any from sqlalchemy.ext.asyncio import ( - AsyncEngine, AsyncSession, async_engine_from_config, async_sessionmaker, - create_async_engine, ) -from syncmaster.server.services.unit_of_work import UnitOfWork from syncmaster.server.settings import DatabaseSettings -from syncmaster.server.settings import ServerAppSettings as Settings - - -def create_engine(connection_uri: str, **engine_kwargs: Any) -> AsyncEngine: - return create_async_engine(url=connection_uri, **engine_kwargs) def create_session_factory(settings: DatabaseSettings) -> async_sessionmaker[AsyncSession]: @@ -27,14 +17,3 @@ def create_session_factory(settings: DatabaseSettings) -> async_sessionmaker[Asy class_=AsyncSession, expire_on_commit=False, ) - - -def get_uow( - session_factory: async_sessionmaker[AsyncSession], - settings: Settings, -) -> Callable[[], AsyncGenerator[UnitOfWork, None]]: - async def wrapper(): - async with session_factory() as session: - yield UnitOfWork(session=session, settings=settings) - - return wrapper diff --git a/syncmaster/server/providers/auth/keycloak_provider.py b/syncmaster/server/providers/auth/keycloak_provider.py index 191f27b4..84be040f 100644 --- a/syncmaster/server/providers/auth/keycloak_provider.py +++ b/syncmaster/server/providers/auth/keycloak_provider.py @@ -1,9 +1,11 @@ # SPDX-FileCopyrightText: 2023-present MTS PJSC # SPDX-License-Identifier: Apache-2.0 import logging +import time from typing import Any, NoReturn from fastapi import FastAPI, Request +from jwcrypto import jwk from jwcrypto.common import JWException from keycloak import KeycloakOpenID, KeycloakOperationError from starlette.middleware.sessions import SessionMiddleware @@ -29,6 +31,8 @@ def __init__(self, settings: KeycloakAuthProviderSettings) -> None: client_secret_key=self.settings.keycloak.client_secret.get_secret_value(), verify=self.settings.keycloak.verify_ssl, ) + self._key: jwk.JWKSet | None = None + self._key_expiration: float = time.monotonic() @classmethod def setup(cls, app: FastAPI) -> FastAPI: @@ -88,7 +92,8 @@ async def get_current_user(self, access_token: str | None, request: Request, uow try: # if user is disabled or blocked in Keycloak after the token is issued, he will # remain authorized until the token expires (not more than 15 minutes in MTS SSO) - token_info = await self.keycloak_openid.a_decode_token(access_token) + key = await self._get_key() + token_info = await self.keycloak_openid.a_decode_token(access_token, key=key) except (KeycloakOperationError, JWException) as e: log.info("Access token is invalid or expired: %s", e) token_info = None @@ -104,8 +109,10 @@ async def get_current_user(self, access_token: str | None, request: Request, uow request.session["access_token"] = new_access_token request.session["refresh_token"] = new_refresh_token + key = await self._get_key() token_info = await self.keycloak_openid.a_decode_token( token=new_access_token, + key=key, ) log.debug("Access token refreshed and decoded successfully.") except (KeycloakOperationError, JWException) as e: @@ -160,3 +167,17 @@ async def logout(self, user: User, refresh_token: str | None) -> None: msg = f"Can't logout user: {user.username}" log.debug("%s. Error: %s", msg, err) raise LogoutError(msg) from err + + async def _get_key(self) -> jwk.JWKSet: + # avoid sending requests to Keycloak for every received token + if self._key is not None and self._key_expiration > time.monotonic(): + return self._key + + key = jwk.JWKSet() + certs = await self.keycloak_openid.a_certs() + for cert in certs["keys"]: + key.add(jwk.JWK(**cert)) + + self._key = key + self._key_expiration = time.monotonic() + self.settings.keycloak.cert_cache_ttl.total_seconds() + return self._key diff --git a/syncmaster/server/settings/auth/keycloak.py b/syncmaster/server/settings/auth/keycloak.py index bd832fb4..f173d4ca 100644 --- a/syncmaster/server/settings/auth/keycloak.py +++ b/syncmaster/server/settings/auth/keycloak.py @@ -1,6 +1,7 @@ # SPDX-FileCopyrightText: 2023-present MTS PJSC # SPDX-License-Identifier: Apache-2.0 import textwrap +from datetime import timedelta from typing import Literal from pydantic import ( @@ -20,6 +21,7 @@ class KeycloakSettings(BaseModel): ui_callback_url: str = Field(description="SyncMaster UI auth callback endpoint") verify_ssl: bool = Field(default=True, description="Verify SSL certificates") scope: str = Field(default="openid", description="Keycloak scope") + cert_cache_ttl: timedelta = Field(default=timedelta(hours=1), description="Keycloak certs cache TTL") class KeycloakCookieSettings(BaseModel): @@ -48,6 +50,7 @@ class KeycloakCookieSettings(BaseModel): same_site: lax https_only: false domain: localhost + cert_cache_ttl: 1D ``` For production environment: @@ -62,6 +65,7 @@ class KeycloakCookieSettings(BaseModel): same_site: strict https_only: true domain: example.com + cert_cache_ttl: 1H ``` """