Inject every models variations and DB model in DB adapters (#84)

* Inject every model variations in router and DB model in DB adapters

* Update documentation and import Tortoise in db module

* Use path operation decorator dependencies for superuser routes
This commit is contained in:
François Voron
2020-01-04 15:36:34 +01:00
committed by GitHub
parent c903b30161
commit 104a6c6bf5
29 changed files with 501 additions and 269 deletions

View File

@@ -2,5 +2,5 @@
__version__ = "0.4.1"
from fastapi_users import models # noqa: F401
from fastapi_users.fastapi_users import FastAPIUsers # noqa: F401
from fastapi_users.models import BaseUser # noqa: F401

View File

@@ -12,3 +12,11 @@ try:
)
except ImportError: # pragma: no cover
pass
try:
from fastapi_users.db.tortoise import ( # noqa: F401
TortoiseBaseUserModel,
TortoiseUserDatabase,
)
except ImportError: # pragma: no cover
pass

View File

@@ -1,41 +1,50 @@
from typing import List, Optional
from typing import Generic, List, Optional, Type
from fastapi.security import OAuth2PasswordRequestForm
from fastapi_users import password
from fastapi_users.models import BaseUserDB
from fastapi_users.models import UD
class BaseUserDatabase:
"""Base adapter for retrieving, creating and updating users from a database."""
class BaseUserDatabase(Generic[UD]):
"""
Base adapter for retrieving, creating and updating users from a database.
async def list(self) -> List[BaseUserDB]:
:param user_db_model: Pydantic model of a DB representation of a user.
"""
user_db_model: Type[UD]
def __init__(self, user_db_model: Type[UD]):
self.user_db_model = user_db_model
async def list(self) -> List[UD]:
"""List all users."""
raise NotImplementedError()
async def get(self, id: str) -> Optional[BaseUserDB]:
async def get(self, id: str) -> Optional[UD]:
"""Get a single user by id."""
raise NotImplementedError()
async def get_by_email(self, email: str) -> Optional[BaseUserDB]:
async def get_by_email(self, email: str) -> Optional[UD]:
"""Get a single user by email."""
raise NotImplementedError()
async def create(self, user: BaseUserDB) -> BaseUserDB:
async def create(self, user: UD) -> UD:
"""Create a user."""
raise NotImplementedError()
async def update(self, user: BaseUserDB) -> BaseUserDB:
async def update(self, user: UD) -> UD:
"""Update a user."""
raise NotImplementedError()
async def delete(self, user: BaseUserDB) -> None:
async def delete(self, user: UD) -> None:
"""Delete a user."""
raise NotImplementedError()
async def authenticate(
self, credentials: OAuth2PasswordRequestForm
) -> Optional[BaseUserDB]:
) -> Optional[UD]:
"""
Authenticate and return a user following an email and a password.

View File

@@ -1,43 +1,45 @@
from typing import List, Optional
from typing import List, Optional, Type
from motor.motor_asyncio import AsyncIOMotorCollection
from fastapi_users.db.base import BaseUserDatabase
from fastapi_users.models import BaseUserDB
from fastapi_users.models import UD
class MongoDBUserDatabase(BaseUserDatabase):
class MongoDBUserDatabase(BaseUserDatabase[UD]):
"""
Database adapter for MongoDB.
:param user_db_model: Pydantic model of a DB representation of a user.
:param collection: Collection instance from `motor`.
"""
collection: AsyncIOMotorCollection
def __init__(self, collection: AsyncIOMotorCollection):
def __init__(self, user_db_model: Type[UD], collection: AsyncIOMotorCollection):
super().__init__(user_db_model)
self.collection = collection
self.collection.create_index("id", unique=True)
self.collection.create_index("email", unique=True)
async def list(self) -> List[BaseUserDB]:
return [BaseUserDB(**user) async for user in self.collection.find()]
async def list(self) -> List[UD]:
return [self.user_db_model(**user) async for user in self.collection.find()]
async def get(self, id: str) -> Optional[BaseUserDB]:
async def get(self, id: str) -> Optional[UD]:
user = await self.collection.find_one({"id": id})
return BaseUserDB(**user) if user else None
return self.user_db_model(**user) if user else None
async def get_by_email(self, email: str) -> Optional[BaseUserDB]:
async def get_by_email(self, email: str) -> Optional[UD]:
user = await self.collection.find_one({"email": email})
return BaseUserDB(**user) if user else None
return self.user_db_model(**user) if user else None
async def create(self, user: BaseUserDB) -> BaseUserDB:
async def create(self, user: UD) -> UD:
await self.collection.insert_one(user.dict())
return user
async def update(self, user: BaseUserDB) -> BaseUserDB:
async def update(self, user: UD) -> UD:
await self.collection.replace_one({"id": user.id}, user.dict())
return user
async def delete(self, user: BaseUserDB) -> None:
async def delete(self, user: UD) -> None:
await self.collection.delete_one({"id": user.id})

View File

@@ -1,10 +1,10 @@
from typing import List, Optional
from typing import List, Optional, Type
from databases import Database
from sqlalchemy import Boolean, Column, String, Table
from fastapi_users.db.base import BaseUserDatabase
from fastapi_users.models import BaseUserDB
from fastapi_users.models import UD
class SQLAlchemyBaseUserTable:
@@ -19,10 +19,11 @@ class SQLAlchemyBaseUserTable:
is_superuser = Column(Boolean, default=False, nullable=False)
class SQLAlchemyUserDatabase(BaseUserDatabase):
class SQLAlchemyUserDatabase(BaseUserDatabase[UD]):
"""
Database adapter for SQLAlchemy.
:param user_db_model: Pydantic model of a DB representation of a user.
:param database: `Database` instance from `encode/databases`.
:param users: SQLAlchemy users table instance.
"""
@@ -30,37 +31,38 @@ class SQLAlchemyUserDatabase(BaseUserDatabase):
database: Database
users: Table
def __init__(self, database: Database, users: Table):
def __init__(self, user_db_model: Type[UD], database: Database, users: Table):
super().__init__(user_db_model)
self.database = database
self.users = users
async def list(self) -> List[BaseUserDB]:
async def list(self) -> List[UD]:
query = self.users.select()
users = await self.database.fetch_all(query)
return [BaseUserDB(**user) for user in users]
return [self.user_db_model(**user) for user in users]
async def get(self, id: str) -> Optional[BaseUserDB]:
async def get(self, id: str) -> Optional[UD]:
query = self.users.select().where(self.users.c.id == id)
user = await self.database.fetch_one(query)
return BaseUserDB(**user) if user else None
return self.user_db_model(**user) if user else None
async def get_by_email(self, email: str) -> Optional[BaseUserDB]:
async def get_by_email(self, email: str) -> Optional[UD]:
query = self.users.select().where(self.users.c.email == email)
user = await self.database.fetch_one(query)
return BaseUserDB(**user) if user else None
return self.user_db_model(**user) if user else None
async def create(self, user: BaseUserDB) -> BaseUserDB:
async def create(self, user: UD) -> UD:
query = self.users.insert().values(**user.dict())
await self.database.execute(query)
return user
async def update(self, user: BaseUserDB) -> BaseUserDB:
async def update(self, user: UD) -> UD:
query = (
self.users.update().where(self.users.c.id == user.id).values(**user.dict())
)
await self.database.execute(query)
return user
async def delete(self, user: BaseUserDB) -> None:
async def delete(self, user: UD) -> None:
query = self.users.delete().where(self.users.c.id == user.id)
await self.database.execute(query)

View File

@@ -3,11 +3,11 @@ from typing import List, Optional, Type
from tortoise import Model, fields
from tortoise.exceptions import DoesNotExist
from fastapi_users.db import BaseUserDatabase
from fastapi_users.models import BaseUserDB
from fastapi_users.db.base import BaseUserDatabase
from fastapi_users.models import UD
class BaseUserModel:
class TortoiseBaseUserModel(Model):
id = fields.CharField(pk=True, generated=False, max_length=255)
email = fields.CharField(index=True, unique=True, null=False, max_length=255)
hashed_password = fields.CharField(null=False, max_length=255)
@@ -15,44 +15,51 @@ class BaseUserModel:
is_superuser = fields.BooleanField(default=False, null=False)
class Meta:
table = "user"
abstract = True
class TortoiseUserDatabase(BaseUserDatabase):
class TortoiseUserDatabase(BaseUserDatabase[UD]):
"""
Database adapter for Tortoise ORM.
model: Type[Model]
:param user_db_model: Pydantic model of a DB representation of a user.
:param model: Tortoise ORM model.
"""
def __init__(self, model: Type[Model]):
model: Type[TortoiseBaseUserModel]
def __init__(self, user_db_model: Type[UD], model: Type[TortoiseBaseUserModel]):
super().__init__(user_db_model)
self.model = model
async def list(self) -> List[BaseUserDB]:
async def list(self) -> List[UD]:
users = await self.model.all()
return [BaseUserDB.from_orm(user) for user in users]
return [self.user_db_model.from_orm(user) for user in users]
async def get(self, id: str) -> Optional[BaseUserDB]:
async def get(self, id: str) -> Optional[UD]:
try:
user = await self.model.get(id=id)
return BaseUserDB.from_orm(user)
return self.user_db_model.from_orm(user)
except DoesNotExist:
return None
async def get_by_email(self, email: str) -> Optional[BaseUserDB]:
async def get_by_email(self, email: str) -> Optional[UD]:
try:
user = await self.model.get(email=email)
return BaseUserDB.from_orm(user)
return self.user_db_model.from_orm(user)
except DoesNotExist:
return None
async def create(self, user: BaseUserDB) -> BaseUserDB:
async def create(self, user: UD) -> UD:
model = self.model(**user.dict())
await model.save()
return user
async def update(self, user: BaseUserDB) -> BaseUserDB:
async def update(self, user: UD) -> UD:
user_dict = user.dict()
user_dict.pop("id") # Tortoise complains if we pass the PK again
await self.model.filter(id=user.id).update(**user_dict)
return user
async def delete(self, user: BaseUserDB) -> None:
async def delete(self, user: UD) -> None:
await self.model.filter(id=user.id).delete()

View File

@@ -1,8 +1,8 @@
from typing import Callable, Sequence, Type
from fastapi_users import models
from fastapi_users.authentication import Authenticator, BaseAuthentication
from fastapi_users.db import BaseUserDatabase
from fastapi_users.models import BaseUser
from fastapi_users.router import Event, UserRouter, get_user_router
@@ -13,6 +13,9 @@ class FastAPIUsers:
:param db: Database adapter instance.
:param auth_backends: List of authentication backends.
:param user_model: Pydantic model of a user.
:param user_create_model: Pydantic model for creating a user.
:param user_update_model: Pydantic model for updating a user.
:param user_db_model: Pydantic model of a DB representation of a user.
:param reset_password_token_secret: Secret to encode reset password token.
:param reset_password_token_lifetime_seconds: Lifetime of reset password token.
@@ -28,7 +31,10 @@ class FastAPIUsers:
self,
db: BaseUserDatabase,
auth_backends: Sequence[BaseAuthentication],
user_model: Type[BaseUser],
user_model: Type[models.BaseUser],
user_create_model: Type[models.BaseUserCreate],
user_update_model: Type[models.BaseUserUpdate],
user_db_model: Type[models.BaseUserDB],
reset_password_token_secret: str,
reset_password_token_lifetime_seconds: int = 3600,
):
@@ -37,6 +43,9 @@ class FastAPIUsers:
self.router = get_user_router(
self.db,
user_model,
user_create_model,
user_update_model,
user_db_model,
self.authenticator,
reset_password_token_secret,
reset_password_token_lifetime_seconds,

View File

@@ -1,5 +1,5 @@
import uuid
from typing import Optional, Type
from typing import Optional, TypeVar
import pydantic
from pydantic import BaseModel, EmailStr
@@ -25,9 +25,6 @@ class BaseUser(BaseModel):
def create_update_dict_superuser(self):
return self.dict(exclude_unset=True, exclude={"id"})
class Config:
orm_mode = True
class BaseUserCreate(BaseUser):
email: EmailStr
@@ -39,23 +36,11 @@ class BaseUserUpdate(BaseUser):
class BaseUserDB(BaseUser):
id: str
hashed_password: str
class Config:
orm_mode = True
class Models:
"""Generate models inheriting from the custom User model."""
def __init__(self, user_model: Type[BaseUser]):
class UserCreate(user_model, BaseUserCreate): # type: ignore
pass
class UserUpdate(user_model, BaseUserUpdate): # type: ignore
pass
class UserDB(user_model, BaseUserDB): # type: ignore
pass
self.User = user_model
self.UserCreate = UserCreate
self.UserUpdate = UserUpdate
self.UserDB = UserDB
UD = TypeVar("UD", bound=BaseUserDB)

View File

@@ -1,7 +1,7 @@
import asyncio
import typing
from collections import defaultdict
from enum import Enum, auto
from typing import Any, Callable, DefaultDict, Dict, List, Type, cast
import jwt
from fastapi import APIRouter, Body, Depends, HTTPException
@@ -10,9 +10,9 @@ from pydantic import EmailStr
from starlette import status
from starlette.responses import Response
from fastapi_users import models
from fastapi_users.authentication import Authenticator, BaseAuthentication
from fastapi_users.db import BaseUserDatabase
from fastapi_users.models import BaseUser, Models
from fastapi_users.password import get_password_hash
from fastapi_users.utils import JWT_ALGORITHM, generate_jwt
@@ -29,13 +29,13 @@ class Event(Enum):
class UserRouter(APIRouter):
event_handlers: typing.DefaultDict[Event, typing.List[typing.Callable]]
event_handlers: DefaultDict[Event, List[Callable]]
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.event_handlers = defaultdict(list)
def add_event_handler(self, event_type: Event, func: typing.Callable) -> None:
def add_event_handler(self, event_type: Event, func: Callable) -> None:
self.event_handlers[event_type].append(func)
async def run_handlers(self, event_type: Event, *args, **kwargs) -> None:
@@ -65,30 +65,30 @@ def _add_login_route(
def get_user_router(
user_db: BaseUserDatabase,
user_model: typing.Type[BaseUser],
user_db: BaseUserDatabase[models.BaseUserDB],
user_model: Type[models.BaseUser],
user_create_model: Type[models.BaseUserCreate],
user_update_model: Type[models.BaseUserUpdate],
user_db_model: Type[models.BaseUserDB],
authenticator: Authenticator,
reset_password_token_secret: str,
reset_password_token_lifetime_seconds: int = 3600,
) -> UserRouter:
"""Generate a router with the authentication routes."""
router = UserRouter()
models = Models(user_model)
reset_password_token_audience = "fastapi-users:reset"
get_current_active_user = authenticator.get_current_active_user
get_current_superuser = authenticator.get_current_superuser
async def _get_or_404(id: str) -> models.UserDB: # type: ignore
async def _get_or_404(id: str) -> models.BaseUserDB:
user = await user_db.get(id)
if user is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND)
return user
async def _update_user(
user: models.UserDB, update_dict: typing.Dict[str, typing.Any] # type: ignore
):
async def _update_user(user: models.BaseUserDB, update_dict: Dict[str, Any]):
for field in update_dict:
if field == "password":
hashed_password = get_password_hash(update_dict[field])
@@ -101,9 +101,10 @@ def get_user_router(
_add_login_route(router, user_db, auth_backend)
@router.post(
"/register", response_model=models.User, status_code=status.HTTP_201_CREATED
"/register", response_model=user_model, status_code=status.HTTP_201_CREATED
)
async def register(user: models.UserCreate): # type: ignore
async def register(user: user_create_model): # type: ignore
user = cast(models.BaseUserCreate, user) # Prevent mypy complain
existing_user = await user_db.get_by_email(user.email)
if existing_user is not None:
@@ -113,7 +114,7 @@ def get_user_router(
)
hashed_password = get_password_hash(user.password)
db_user = models.UserDB(
db_user = user_db_model(
**user.create_update_dict(), hashed_password=hashed_password
)
created_user = await user_db.create(db_user)
@@ -168,48 +169,60 @@ def get_user_router(
detail=ErrorCode.RESET_PASSWORD_BAD_TOKEN,
)
@router.get("/me", response_model=models.User)
@router.get("/me", response_model=user_model)
async def me(
user: models.UserDB = Depends(get_current_active_user), # type: ignore
user: user_db_model = Depends(get_current_active_user), # type: ignore
):
return user
@router.patch("/me", response_model=models.User)
@router.patch("/me", response_model=user_model)
async def update_me(
updated_user: models.UserUpdate, # type: ignore
user: models.UserDB = Depends(get_current_active_user), # type: ignore
updated_user: user_update_model, # type: ignore
user: user_db_model = Depends(get_current_active_user), # type: ignore
):
updated_user = cast(
models.BaseUserUpdate, updated_user,
) # Prevent mypy complain
updated_user_data = updated_user.create_update_dict()
return await _update_user(user, updated_user_data)
@router.get("/", response_model=typing.List[models.User]) # type: ignore
async def list_users(
superuser: models.UserDB = Depends(get_current_superuser), # type: ignore
):
@router.get(
"/",
response_model=List[user_model], # type: ignore
dependencies=[Depends(get_current_superuser)],
)
async def list_users():
return await user_db.list()
@router.get("/{id}", response_model=models.User)
async def get_user(
id: str,
superuser: models.UserDB = Depends(get_current_superuser), # type: ignore
):
@router.get(
"/{id}",
response_model=user_model,
dependencies=[Depends(get_current_superuser)],
)
async def get_user(id: str,):
return await _get_or_404(id)
@router.patch("/{id}", response_model=models.User)
@router.patch(
"/{id}",
response_model=user_model,
dependencies=[Depends(get_current_superuser)],
)
async def update_user(
id: str,
updated_user: models.UserUpdate, # type: ignore
superuser: models.UserDB = Depends(get_current_superuser), # type: ignore
id: str, updated_user: user_update_model, # type: ignore
):
updated_user = cast(
models.BaseUserUpdate, updated_user,
) # Prevent mypy complain
user = await _get_or_404(id)
updated_user_data = updated_user.create_update_dict_superuser()
return await _update_user(user, updated_user_data)
@router.delete("/{id}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_user(
id: str,
superuser: models.UserDB = Depends(get_current_superuser), # type: ignore
):
@router.delete(
"/{id}",
status_code=status.HTTP_204_NO_CONTENT,
dependencies=[Depends(get_current_superuser)],
)
async def delete_user(id: str):
user = await _get_or_404(id)
await user_db.delete(user)
return None