Update users to use the new CRUD types
This commit is contained in:
parent
300cb7bd67
commit
bdf74c2dbd
3 changed files with 37 additions and 56 deletions
|
|
@ -6,7 +6,7 @@ from src.config.email import EmailAddress
|
||||||
from sqlalchemy import String, VARCHAR, Engine, select
|
from sqlalchemy import String, VARCHAR, Engine, select
|
||||||
from sqlalchemy.orm import mapped_column, Mapped, Session, relationship
|
from sqlalchemy.orm import mapped_column, Mapped, Session, relationship
|
||||||
from argon2 import PasswordHasher
|
from argon2 import PasswordHasher
|
||||||
from src.infra.db import Base
|
from src.infra.db import Base, CRUDRepo
|
||||||
from src.services.users.data import UserDTO, User
|
from src.services.users.data import UserDTO, User
|
||||||
from src.services.users.repo import UserRepo
|
from src.services.users.repo import UserRepo
|
||||||
|
|
||||||
|
|
@ -33,42 +33,7 @@ def model_to_user(db_user: UserModel) -> User:
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class UserRepoImpl(UserRepo):
|
def mutate_user(db_user: UserModel, user: User):
|
||||||
def __init__(self, db: Engine):
|
|
||||||
self._db = db
|
|
||||||
self._hasher = PasswordHasher()
|
|
||||||
self._logger = logging.getLogger('UserRepoImpl')
|
|
||||||
|
|
||||||
def create_user(self, user: UserDTO) -> User:
|
|
||||||
pw = self._hasher.hash(user.raw_password.expose_secret())
|
|
||||||
db_user = UserModel(email=user.email, password_hash=pw)
|
|
||||||
with Session(self._db) as session:
|
|
||||||
session.add(db_user)
|
|
||||||
session.commit()
|
|
||||||
|
|
||||||
return model_to_user(db_user)
|
|
||||||
|
|
||||||
def get_user_by_id(self, user_id: int) -> User | None:
|
|
||||||
with Session(self._db) as session:
|
|
||||||
user = session.get(UserModel, user_id)
|
|
||||||
|
|
||||||
if user:
|
|
||||||
return model_to_user(user)
|
|
||||||
|
|
||||||
def get_user_by_email(self, email: EmailAddress) -> User | None:
|
|
||||||
with Session(self._db) as session:
|
|
||||||
user = session.scalar(select(UserModel).where(UserModel.email == email))
|
|
||||||
|
|
||||||
if user:
|
|
||||||
return model_to_user(user)
|
|
||||||
|
|
||||||
def update_user(self, user: User):
|
|
||||||
with Session(self._db) as session:
|
|
||||||
db_user = session.get(UserModel, user.id)
|
|
||||||
if db_user is None:
|
|
||||||
self._logger.warning(f'Attempted to update non-existing user {user.id}')
|
|
||||||
return
|
|
||||||
|
|
||||||
# These fields can be set directly
|
# These fields can be set directly
|
||||||
db_user.email_confirmed = user.email_confirmed
|
db_user.email_confirmed = user.email_confirmed
|
||||||
db_user.password_hash = user.password_hash.expose_secret()
|
db_user.password_hash = user.password_hash.expose_secret()
|
||||||
|
|
@ -78,7 +43,28 @@ class UserRepoImpl(UserRepo):
|
||||||
db_user.email = user.email
|
db_user.email = user.email
|
||||||
db_user.email_confirmed = False
|
db_user.email_confirmed = False
|
||||||
|
|
||||||
session.commit()
|
|
||||||
|
def make_create_user(hasher: PasswordHasher):
|
||||||
|
def create_user(user: UserDTO) -> UserModel:
|
||||||
|
pw = hasher.hash(user.raw_password.expose_secret())
|
||||||
|
return UserModel(email=user.email, password_hash=pw)
|
||||||
|
|
||||||
|
return create_user
|
||||||
|
|
||||||
|
|
||||||
|
class UserRepoImpl(CRUDRepo[UserModel, User, UserDTO], UserRepo):
|
||||||
|
def __init__(self, db: Engine):
|
||||||
|
hasher = PasswordHasher()
|
||||||
|
super().__init__(db, UserModel, model_to_user, mutate_user, make_create_user(hasher))
|
||||||
|
self._logger = logging.getLogger('UserRepoImpl')
|
||||||
|
self._hasher = hasher
|
||||||
|
|
||||||
|
def get_user_by_email(self, email: EmailAddress) -> User | None:
|
||||||
|
with Session(self._db) as session:
|
||||||
|
user = session.scalar(select(UserModel).where(UserModel.email == email))
|
||||||
|
|
||||||
|
if user:
|
||||||
|
return model_to_user(user)
|
||||||
|
|
||||||
def auth_as_user(self, user: UserDTO) -> User | None:
|
def auth_as_user(self, user: UserDTO) -> User | None:
|
||||||
full_user = self.get_user_by_email(user.email)
|
full_user = self.get_user_by_email(user.email)
|
||||||
|
|
|
||||||
|
|
@ -1,26 +1,15 @@
|
||||||
|
from src.infra.db import CRUD
|
||||||
import abc
|
import abc
|
||||||
from src.config.email import EmailAddress
|
from src.config.email import EmailAddress
|
||||||
|
|
||||||
from .data import User, UserDTO
|
from .data import User, UserDTO
|
||||||
|
|
||||||
|
|
||||||
class UserRepo(abc.ABC):
|
class UserRepo(CRUD[User, UserDTO]):
|
||||||
@abc.abstractmethod
|
|
||||||
def create_user(self, user: UserDTO) -> User:
|
|
||||||
pass
|
|
||||||
|
|
||||||
@abc.abstractmethod
|
|
||||||
def get_user_by_id(self, user_id: int) -> User | None:
|
|
||||||
pass
|
|
||||||
|
|
||||||
@abc.abstractmethod
|
@abc.abstractmethod
|
||||||
def get_user_by_email(self, email: EmailAddress) -> User | None:
|
def get_user_by_email(self, email: EmailAddress) -> User | None:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@abc.abstractmethod
|
|
||||||
def update_user(self, user: User):
|
|
||||||
pass
|
|
||||||
|
|
||||||
@abc.abstractmethod
|
@abc.abstractmethod
|
||||||
def auth_as_user(self, user: UserDTO) -> User | None:
|
def auth_as_user(self, user: UserDTO) -> User | None:
|
||||||
pass
|
pass
|
||||||
|
|
|
||||||
|
|
@ -65,14 +65,20 @@ class UserService:
|
||||||
if not claim or claim.sub != user_id:
|
if not claim or claim.sub != user_id:
|
||||||
raise EmailConfirmationTokenInvalid
|
raise EmailConfirmationTokenInvalid
|
||||||
|
|
||||||
user = self._repo.get_user_by_id(user_id)
|
user = self._repo.get(user_id)
|
||||||
# If there is no user or the users emails don't match or the user is already confirmed then consider
|
# If there is no user or the users emails don't match or the user is already confirmed then consider
|
||||||
# the token invalid for this request.
|
# the token invalid for this request.
|
||||||
if not user or user.email != claim.email or user.email_confirmed:
|
if not user or user.email != claim.email or user.email_confirmed:
|
||||||
raise EmailConfirmationTokenInvalid
|
raise EmailConfirmationTokenInvalid
|
||||||
|
|
||||||
|
def update(user: User) -> User:
|
||||||
user.email_confirmed = True
|
user.email_confirmed = True
|
||||||
self._repo.update_user(user)
|
return user
|
||||||
|
|
||||||
|
self._repo.update(
|
||||||
|
user_id,
|
||||||
|
update,
|
||||||
|
)
|
||||||
|
|
||||||
def _send_confirmation_email(self, user: User):
|
def _send_confirmation_email(self, user: User):
|
||||||
env = Environment(loader=PackageLoader('src'), autoescape=select_autoescape())
|
env = Environment(loader=PackageLoader('src'), autoescape=select_autoescape())
|
||||||
|
|
@ -96,13 +102,13 @@ class UserService:
|
||||||
raise SignupError('The given password was not acceptable')
|
raise SignupError('The given password was not acceptable')
|
||||||
|
|
||||||
# Create a user in persistence
|
# Create a user in persistence
|
||||||
created_user = self._repo.create_user(user)
|
created_user = self._repo.create(user, None)
|
||||||
# send a confirmation email
|
# send a confirmation email
|
||||||
self._send_confirmation_email(created_user)
|
self._send_confirmation_email(created_user)
|
||||||
|
|
||||||
return created_user.to_profile()
|
return created_user.to_profile()
|
||||||
|
|
||||||
def get_user_by_id(self, user_id: int) -> UserProfile | None:
|
def get_user_by_id(self, user_id: int) -> UserProfile | None:
|
||||||
user = self._repo.get_user_by_id(user_id)
|
user = self._repo.get(user_id)
|
||||||
if user:
|
if user:
|
||||||
return user.to_profile()
|
return user.to_profile()
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue