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/409.improvement.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Drop FastAPI `dependency_overrides` feature. See [#14441](https://github.com/fastapi/fastapi/pull/14441).
5 changes: 4 additions & 1 deletion syncmaster/db/factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,19 +6,22 @@
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(engine: AsyncEngine) -> async_sessionmaker[AsyncSession]:
def create_session_factory(settings: DatabaseSettings) -> async_sessionmaker[AsyncSession]:
engine = async_engine_from_config(settings.model_dump(), prefix="")
return async_sessionmaker(
bind=engine,
class_=AsyncSession,
Expand Down
23 changes: 3 additions & 20 deletions syncmaster/server/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,10 +6,9 @@
from fastapi import FastAPI, HTTPException
from fastapi.exceptions import RequestValidationError
from pydantic import ValidationError
from sqlalchemy.ext.asyncio import async_engine_from_config

from syncmaster import _raw_version as syncmaster_version
from syncmaster.db.factory import create_session_factory, get_uow
from syncmaster.db.factory import create_session_factory
from syncmaster.exceptions import SyncmasterError
from syncmaster.server.api.router import api_router
from syncmaster.server.handler import (
Expand All @@ -19,7 +18,6 @@
validation_exception_handler,
)
from syncmaster.server.middlewares import apply_middlewares
from syncmaster.server.services.unit_of_work import UnitOfWork
from syncmaster.server.settings import ServerAppSettings as Settings
from syncmaster.settings.logging import setup_logging

Expand Down Expand Up @@ -49,30 +47,15 @@ def application_factory(settings: Settings) -> FastAPI:
)
application.state.settings = settings
application.state.celery = celery_factory(settings)
application.state.session_factory = create_session_factory(settings.database)

application.include_router(api_router)
application.exception_handler(RequestValidationError)(validation_exception_handler)
application.exception_handler(ValidationError)(validation_exception_handler)
application.exception_handler(SyncmasterError)(syncmsater_exception_handler)
application.exception_handler(HTTPException)(http_exception_handler)
application.exception_handler(Exception)(unknown_exception_handler)

engine = async_engine_from_config(settings.database.model_dump(), prefix="")
session_factory = create_session_factory(engine=engine)

async def get_settings():
return settings

async def get_celery():
return application.state.celery

application.dependency_overrides.update(
{
Settings: get_settings,
UnitOfWork: get_uow(session_factory, settings=settings),
Celery: get_celery,
},
)

auth_class: type[AuthProvider] = settings.auth.provider # type: ignore[assignment]
auth_class.setup(application)

Expand Down
18 changes: 9 additions & 9 deletions syncmaster/server/api/v1/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,13 +11,10 @@
from syncmaster.errors.schemas.invalid_request import InvalidRequestSchema
from syncmaster.errors.schemas.not_authorized import NotAuthorizedSchema
from syncmaster.schemas.v1.auth import AuthTokenSchema
from syncmaster.server.dependencies import Stub
from syncmaster.server.providers.auth import (
AuthProvider,
DummyAuthProvider,
KeycloakAuthProvider,
)
from syncmaster.server.providers.auth import AuthProvider
from syncmaster.server.services.auth import get_auth_provider
from syncmaster.server.services.get_user import get_user
from syncmaster.server.services.unit_of_work import UnitOfWork

router = APIRouter(
prefix="/auth",
Expand All @@ -28,10 +25,12 @@

@router.post("/token")
async def token(
auth_provider: Annotated[DummyAuthProvider, Depends(Stub(AuthProvider))],
uow: Annotated[UnitOfWork, Depends()],
form_data: Annotated[OAuth2PasswordRequestForm, Depends()],
auth_provider: Annotated[AuthProvider, Depends(get_auth_provider)],
) -> AuthTokenSchema:
token = await auth_provider.get_token_password_grant(
uow=uow,
grant_type=form_data.grant_type,
login=form_data.username,
password=form_data.password,
Expand All @@ -46,10 +45,11 @@ async def token(
async def auth_callback(
request: Request,
code: str,
auth_provider: Annotated[KeycloakAuthProvider, Depends(Stub(AuthProvider))],
auth_provider: Annotated[AuthProvider, Depends(get_auth_provider)],
):
token = await auth_provider.get_token_authorization_code_grant(
code=code,
request=request,
)
request.session["access_token"] = token["access_token"]
request.session["refresh_token"] = token["refresh_token"]
Expand All @@ -64,7 +64,7 @@ async def auth_callback(
async def logout(
request: Request,
current_user: Annotated[User, Depends(get_user())],
auth_provider: Annotated[KeycloakAuthProvider, Depends(Stub(AuthProvider))],
auth_provider: Annotated[AuthProvider, Depends(get_auth_provider)],
):
refresh_token = request.session.get("refresh_token", None)
request.session.clear()
Expand Down
9 changes: 6 additions & 3 deletions syncmaster/server/api/v1/runs.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
from typing import Annotated

from celery import Celery
from fastapi import APIRouter, Depends, Query
from fastapi import APIRouter, Depends, Query, Request
from kombu.exceptions import KombuError

from syncmaster.db.models import RunType, Status, User
Expand All @@ -20,13 +20,16 @@
ReadRunSchema,
RunPageSchema,
)
from syncmaster.server.dependencies import Stub
from syncmaster.server.services.get_user import get_user
from syncmaster.server.services.unit_of_work import UnitOfWork

router = APIRouter(tags=["Runs"], responses=get_error_responses())


async def get_celery(request: Request) -> Celery:
return request.app.state.celery


@router.get("/runs")
async def read_runs( # noqa: PLR0913, PLR0917
transfer_id: int,
Expand Down Expand Up @@ -81,7 +84,7 @@ async def read_run(
@router.post("/runs")
async def start_run(
create_run_data: CreateRunSchema,
celery: Annotated[Celery, Depends(Stub(Celery))],
celery: Annotated[Celery, Depends(get_celery)],
unit_of_work: Annotated[UnitOfWork, Depends(UnitOfWork)],
current_user: Annotated[User, Depends(get_user())],
) -> ReadRunSchema:
Expand Down
2 changes: 0 additions & 2 deletions syncmaster/server/dependencies/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,7 @@
# SPDX-License-Identifier: Apache-2.0

from syncmaster.server.dependencies.get_access_token import get_access_token
from syncmaster.server.dependencies.stub import Stub

__all__ = [
"Stub",
"get_access_token",
]
49 changes: 0 additions & 49 deletions syncmaster/server/dependencies/stub.py

This file was deleted.

18 changes: 9 additions & 9 deletions syncmaster/server/providers/auth/base_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from fastapi import FastAPI, Request

from syncmaster.db.models import User
from syncmaster.server.services.unit_of_work import UnitOfWork


class AuthProvider(ABC):
Expand All @@ -22,7 +23,7 @@ def setup(cls, app: FastAPI) -> FastAPI:
"""
This method is called by `application_factory`.

Here you should add dependency overrides for auth provider,
Here you should configure your auth provider, set `app.state.auth_provider`
and return new `app` object.

Examples
Expand All @@ -31,28 +32,25 @@ def setup(cls, app: FastAPI) -> FastAPI:
```python
from fastapi import FastAPI
from my_awesome_auth_provider.settings import MyAwesomeAuthProviderSettings
from syncmaster.server.dependencies import Stub

class MyAwesomeAuthProvider(AuthProvider):
def setup(app):
app.dependency_overrides[AuthProvider] = MyAwesomeAuthProvider

# `settings_object_factory` returns MyAwesomeAuthProviderSettings object
app.dependency_overrides[MyAwesomeAuthProviderSettings] = settings_object_factory
settings_dict = app.state.settings.auth.model_dump(exclude={"provider})
settings = MyAwesomeAuthProviderSettings.model_validate(settings_dict)
app.state.auth_provider = MyAwesomeAuthProvider(settings)
return app

def __init__(
self,
settings: Annotated[MyAwesomeAuthProviderSettings, Depends(Stub(MyAwesomeAuthProviderSettings))],
settings: MyAwesomeAuthProviderSettings,
):
# settings object is set automatically by FastAPI's dependency_overrides
self.settings = settings
```
"""
...

@abstractmethod
async def get_current_user(self, access_token: str | None, request: Request) -> User:
async def get_current_user(self, access_token: str | None, request: Request, uow: UnitOfWork) -> User:
"""
This method should return currently logged in user.

Expand All @@ -71,6 +69,7 @@ async def get_current_user(self, access_token: str | None, request: Request) ->
@abstractmethod
async def get_token_password_grant( # noqa: PLR0913, PLR0917
self,
uow: UnitOfWork,
grant_type: str | None = None,
login: str | None = None,
password: str | None = None,
Expand Down Expand Up @@ -103,6 +102,7 @@ async def get_token_password_grant( # noqa: PLR0913, PLR0917
async def get_token_authorization_code_grant(
self,
code: str,
request: Request,
scopes: list[str] | None = None,
client_id: str | None = None,
client_secret: str | None = None,
Expand Down
32 changes: 13 additions & 19 deletions syncmaster/server/providers/auth/dummy_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,14 +4,13 @@
import logging
from pprint import pformat
from time import time
from typing import Annotated, Any
from typing import Any

from fastapi import Depends, FastAPI
from fastapi import FastAPI, Request

from syncmaster.db.models import User
from syncmaster.exceptions import EntityNotFoundError
from syncmaster.exceptions.auth import AuthorizationError
from syncmaster.server.dependencies import Stub
from syncmaster.server.providers.auth.base_provider import AuthProvider
from syncmaster.server.services.unit_of_work import UnitOfWork
from syncmaster.server.settings.auth.dummy import DummyAuthProviderSettings
Expand All @@ -21,36 +20,30 @@


class DummyAuthProvider(AuthProvider):
def __init__(
self,
settings: Annotated[DummyAuthProviderSettings, Depends(Stub(DummyAuthProviderSettings))],
unit_of_work: Annotated[UnitOfWork, Depends()],
) -> None:
def __init__(self, settings: DummyAuthProviderSettings) -> None:
self._settings = settings
self._uow = unit_of_work

@classmethod
def setup(cls, app: FastAPI) -> FastAPI:
settings = DummyAuthProviderSettings.model_validate(app.state.settings.auth.model_dump(exclude={"provider"}))
log.info("Using %s provider with settings:\n%s", cls.__name__, pformat(settings))

async def get_settings():
return settings

app.dependency_overrides[AuthProvider] = cls
app.dependency_overrides[DummyAuthProviderSettings] = get_settings
app.state.auth_provider = cls(settings=settings)
return app

async def get_current_user(self, access_token: str | None, *args, **kwargs) -> User:
async def get_current_user(
self, access_token: str | None, request: Request, uow: UnitOfWork, *args, **kwargs
) -> User:
if not access_token:
msg = "Missing auth credentials"
raise AuthorizationError(msg)

user_id = self._get_user_id_from_token(access_token)
return await self._uow.user.read_by_id(user_id)
return await uow.user.read_by_id(user_id)

async def get_token_password_grant( # noqa: PLR0913, PLR0917
self,
uow: UnitOfWork,
grant_type: str | None = None,
login: str | None = None,
password: str | None = None,
Expand All @@ -63,11 +56,11 @@ async def get_token_password_grant( # noqa: PLR0913, PLR0917
raise AuthorizationError(msg)

log.info("Get/create user %r in database", login)
async with self._uow:
async with uow:
try:
user = await self._uow.user.read_by_username(login)
user = await uow.user.read_by_username(login)
except EntityNotFoundError:
user = await self._uow.user.create(username=login)
user = await uow.user.create(username=login)

log.info("User with id %r found", user.id)
if not user.is_active:
Expand Down Expand Up @@ -110,6 +103,7 @@ def _get_user_id_from_token(self, token: str) -> int:
async def get_token_authorization_code_grant(
self,
code: str,
request: Request,
scopes: list[str] | None = None,
client_id: str | None = None,
client_secret: str | None = None,
Expand Down
Loading