Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions mddocs/docs/changelog/next_release/410.improvement.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Cache Keycloak certs for JWT token validation. Previously every API request with KeycloakAuthProvider enabled made a request to Keycloak.
21 changes: 0 additions & 21 deletions syncmaster/db/factory.py
Original file line number Diff line number Diff line change
@@ -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]:
Expand All @@ -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
23 changes: 22 additions & 1 deletion syncmaster/server/providers/auth/keycloak_provider.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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
4 changes: 4 additions & 0 deletions syncmaster/server/settings/auth/keycloak.py
Original file line number Diff line number Diff line change
@@ -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 (
Expand All @@ -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):
Expand Down Expand Up @@ -48,6 +50,7 @@ class KeycloakCookieSettings(BaseModel):
same_site: lax
https_only: false
domain: localhost
cert_cache_ttl: 1D
```
For production environment:

Expand All @@ -62,6 +65,7 @@ class KeycloakCookieSettings(BaseModel):
same_site: strict
https_only: true
domain: example.com
cert_cache_ttl: 1H
```
"""

Expand Down