diff --git a/src/infra/db.py b/src/infra/db.py index 6359df4..841cbb8 100644 --- a/src/infra/db.py +++ b/src/infra/db.py @@ -1,8 +1,6 @@ -from typing import Callable -from sqlalchemy.orm import DeclarativeBase, Session +from sqlalchemy.orm import DeclarativeBase from src.config.database import Database as DatabaseConfig -from sqlalchemy import create_engine, Engine -import abc +from sqlalchemy import create_engine class Base(DeclarativeBase): @@ -15,81 +13,3 @@ def get_database(config: DatabaseConfig): Base.metadata.create_all(engine) return engine - - -class UpdateMissingEntryError(Exception): - def __init__(self, model_name: str, id: int): - super().__init__(f'Attempted to update non-existing {model_name} {id}') - - -class CRUD[T, C](abc.ABC): - @abc.abstractmethod - def create(self, domain_item: C, verify: Callable[[T], None] | None) -> T: - pass - - @abc.abstractmethod - def get(self, item_id: int) -> T | None: - pass - - @abc.abstractmethod - def update(self, item_id: int, transform: Callable[[T], T]) -> T: - pass - - @abc.abstractmethod - def delete(self, item_id: int, verify: Callable[[T], None] | None): - pass - - -class CRUDRepo[T: Base, U, C](CRUD[U, C]): - def __init__( - self, - db: Engine, - root: type[T], - model_to_domain: Callable[[T], U], - mutate_from_domain: Callable[[T, U], None], - create_row: Callable[[C], T], - ): - self._db = db - self._root = root - self._model_to_domain = model_to_domain - self._mutate_from_domain = mutate_from_domain - self._create_row = create_row - - def get(self, item_id: int) -> U | None: - with Session(self._db) as session: - db_item = session.get(self._root, item_id) - if db_item: - return self._model_to_domain(db_item) - - def create(self, domain_item: C, verify: Callable[[U], None] | None) -> U: - item = self._create_row(domain_item) - - with Session(self._db) as session: - session.add(item) - session.flush() - to_create = self._model_to_domain(item) - if verify: - verify(to_create) - session.commit() - return to_create - - def update(self, item_id: int, transform: Callable[[U], U]) -> U: - with Session(self._db) as session: - db_item = session.get(self._root, item_id) - if db_item is None: - raise UpdateMissingEntryError(self._root.__tablename__, item_id) - - domain_item = self._model_to_domain(db_item) - updated = transform(domain_item) - self._mutate_from_domain(db_item, updated) - session.commit() - return updated - - def delete(self, item_id: int, verify: Callable[[U], None] | None): - with Session(self._db) as session: - db_item = session.get(self._root, item_id) - if db_item: - if verify: - verify(self._model_to_domain(db_item)) - session.delete(db_item) - session.commit() diff --git a/src/infra/subscription.py b/src/infra/subscription.py index 5575100..40da20f 100644 --- a/src/infra/subscription.py +++ b/src/infra/subscription.py @@ -1,17 +1,15 @@ -from dataclasses import asdict -import logging from src.infra.publication import model_to_order from typing import TYPE_CHECKING from datetime import datetime -from src.infra.users import UserModel +from src.infra.users import UserModel, model_to_user from sqlalchemy.orm import Mapped, mapped_column, relationship, Session, joinedload -from src.infra.db import Base, CRUDRepo -from sqlalchemy import Engine, ForeignKey, CheckConstraint, DateTime, func, UniqueConstraint, select, delete +from src.infra.db import Base +from sqlalchemy import Engine, ForeignKey, CheckConstraint, DateTime, func, UniqueConstraint, select from src.services.subscription import ( SubscriptionRepo, Subscription, SubscriptionId, - SubscriptionCreateParams, + UpdateSubscription, AvailableSubscriptionEntry, ) @@ -48,24 +46,16 @@ class SubscriptionModel(Base): def model_to_subscription(db_sub: SubscriptionModel) -> Subscription: return Subscription( id=db_sub.id, - user_id=db_sub.user_id, + user=model_to_user(db_sub.user).to_profile(), publication_order=model_to_order(db_sub.order).to_profile(), sequence_seen=db_sub.sequence_seen, start=db_sub.start, ) -def mutate_subscription(db_sub: SubscriptionModel, sub: Subscription): - pass - - -class SubscriptionRepoImpl(SubscriptionRepo, CRUDRepo[SubscriptionModel, Subscription, SubscriptionCreateParams]): +class SubscriptionRepoImpl(SubscriptionRepo): def __init__(self, db: Engine): - super().__init__( - db, SubscriptionModel, model_to_subscription, mutate_subscription, lambda u: SubscriptionModel(**asdict(u)) - ) - - self._logger = logging.getLogger('UserRepoImpl') + self._db = db def get_subscriptions_for_user(self, user_id: int) -> list[Subscription]: with Session(self._db) as session: @@ -75,6 +65,7 @@ class SubscriptionRepoImpl(SubscriptionRepo, CRUDRepo[SubscriptionModel, Subscri .where(SubscriptionModel.user_id == user_id) .options( joinedload(SubscriptionModel.order), + joinedload(SubscriptionModel.user), ) ) .unique() @@ -87,12 +78,21 @@ class SubscriptionRepoImpl(SubscriptionRepo, CRUDRepo[SubscriptionModel, Subscri db_sub = session.get( SubscriptionModel, sub_id, - options=[joinedload(SubscriptionModel.order)], + options=[joinedload(SubscriptionModel.order), joinedload(SubscriptionModel.user)], ) if db_sub: return model_to_subscription(db_sub) + def create_subscription_for_user(self, user_id: int, order_id: int) -> Subscription: + pass + + def delete_subscription(self, subscription_id: SubscriptionId): + pass + + def update_subscription(self, update: UpdateSubscription) -> Subscription: + pass + def get_available_entries(self, user_id: int) -> dict[SubscriptionId, list[AvailableSubscriptionEntry]]: return {} diff --git a/src/infra/users.py b/src/infra/users.py index a01e4bc..f358acb 100644 --- a/src/infra/users.py +++ b/src/infra/users.py @@ -6,7 +6,7 @@ from src.config.email import EmailAddress from sqlalchemy import String, VARCHAR, Engine, select from sqlalchemy.orm import mapped_column, Mapped, Session, relationship from argon2 import PasswordHasher -from src.infra.db import Base, CRUDRepo +from src.infra.db import Base from src.services.users.data import UserDTO, User from src.services.users.repo import UserRepo @@ -33,31 +33,27 @@ def model_to_user(db_user: UserModel) -> User: ) -def mutate_user(db_user: UserModel, user: User): - # These fields can be set directly - db_user.email_confirmed = user.email_confirmed - db_user.password_hash = user.password_hash.expose_secret() - - # If the email is being updated then it should not be considered confirmed. - if db_user.email != user.email: - db_user.email = user.email - db_user.email_confirmed = False - - -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): +class UserRepoImpl(UserRepo): def __init__(self, db: Engine): - hasher = PasswordHasher() - super().__init__(db, UserModel, model_to_user, mutate_user, make_create_user(hasher)) + self._db = db + self._hasher = PasswordHasher() self._logger = logging.getLogger('UserRepoImpl') - self._hasher = hasher + + 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: @@ -66,6 +62,24 @@ class UserRepoImpl(CRUDRepo[UserModel, User, UserDTO], UserRepo): 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 + db_user.email_confirmed = user.email_confirmed + db_user.password_hash = user.password_hash.expose_secret() + + # If the email is being updated then it should not be considered confirmed. + if db_user.email != user.email: + db_user.email = user.email + db_user.email_confirmed = False + + session.commit() + def auth_as_user(self, user: UserDTO) -> User | None: full_user = self.get_user_by_email(user.email) if full_user is None: diff --git a/src/services/subscription.py b/src/services/subscription.py index 5eb5d51..25f70a7 100644 --- a/src/services/subscription.py +++ b/src/services/subscription.py @@ -1,4 +1,3 @@ -from src.infra.db import CRUD import abc from datetime import datetime from src.services.publications import ( @@ -6,6 +5,7 @@ from src.services.publications import ( Sequence as PublicationSequence, Publication, ) +from src.services.users.data import UserProfile from dataclasses import dataclass SubscriptionId = int @@ -20,7 +20,7 @@ class SubscriptionCreateParams: @dataclass class Subscription: id: SubscriptionId - user_id: int + user: UserProfile publication_order: PublicationOrderProfile sequence_seen: int start: datetime @@ -29,6 +29,8 @@ class Subscription: @dataclass class UpdateSubscription: id: int + user_id: int + publication_order_id: int | None = None sequence_seen: int | None = None start: datetime | None = None @@ -41,11 +43,27 @@ class AvailableSubscriptionEntry: sequence: PublicationSequence -class SubscriptionRepo(CRUD[Subscription, SubscriptionCreateParams]): +class SubscriptionRepo(abc.ABC): @abc.abstractmethod def get_subscriptions_for_user(self, user_id: int) -> list[Subscription]: pass + @abc.abstractmethod + def get_subscription_by_id(self, sub_id: SubscriptionId) -> Subscription | None: + pass + + @abc.abstractmethod + def create_subscription_for_user(self, user_id: int, order_id: int) -> Subscription: + pass + + @abc.abstractmethod + def delete_subscription(self, subscription_id: SubscriptionId): + pass + + @abc.abstractmethod + def update_subscription(self, subscription: UpdateSubscription) -> Subscription: + pass + @abc.abstractmethod def get_available_entries(self, user_id: int) -> dict[SubscriptionId, list[AvailableSubscriptionEntry]]: pass @@ -73,11 +91,15 @@ class SubscriptionService: return self._repo.get_available_entries(user_id) def create_suscription(self, params: SubscriptionCreateParams) -> Subscription: - return self._repo.create(params, None) + return self._repo.create_subscription_for_user(params.user_id, params.order_id) def delete_subcription(self, user_id: int, subscription_id: SubscriptionId): - def verify(subscription: Subscription): - if subscription.user_id != user_id: - raise ValueError('Only the owning user can delete a subscription') + subscription = self._repo.get_subscription_by_id(subscription_id) + if subscription is None: + # TODO Domain error types + raise RuntimeError('not found') + if subscription.user.id != user_id: + # TODO: Domain error types + raise RuntimeError('Cannot delete a subscription for another user') - self._repo.delete(subscription_id, verify) + self._repo.delete_subscription(subscription_id) diff --git a/src/services/users/repo.py b/src/services/users/repo.py index 6f8f40c..f50836b 100644 --- a/src/services/users/repo.py +++ b/src/services/users/repo.py @@ -1,15 +1,26 @@ -from src.infra.db import CRUD import abc from src.config.email import EmailAddress from .data import User, UserDTO -class UserRepo(CRUD[User, UserDTO]): +class UserRepo(abc.ABC): + @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 def get_user_by_email(self, email: EmailAddress) -> User | None: pass + @abc.abstractmethod + def update_user(self, user: User): + pass + @abc.abstractmethod def auth_as_user(self, user: UserDTO) -> User | None: pass diff --git a/src/services/users/service.py b/src/services/users/service.py index ca916b6..d481cc0 100644 --- a/src/services/users/service.py +++ b/src/services/users/service.py @@ -65,20 +65,14 @@ class UserService: if not claim or claim.sub != user_id: raise EmailConfirmationTokenInvalid - user = self._repo.get(user_id) + user = self._repo.get_user_by_id(user_id) # 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. if not user or user.email != claim.email or user.email_confirmed: raise EmailConfirmationTokenInvalid - def update(user: User) -> User: - user.email_confirmed = True - return user - - self._repo.update( - user_id, - update, - ) + user.email_confirmed = True + self._repo.update_user(user) def _send_confirmation_email(self, user: User): env = Environment(loader=PackageLoader('src'), autoescape=select_autoescape()) @@ -102,13 +96,13 @@ class UserService: raise SignupError('The given password was not acceptable') # Create a user in persistence - created_user = self._repo.create(user, None) + created_user = self._repo.create_user(user) # send a confirmation email self._send_confirmation_email(created_user) return created_user.to_profile() def get_user_by_id(self, user_id: int) -> UserProfile | None: - user = self._repo.get(user_id) + user = self._repo.get_user_by_id(user_id) if user: return user.to_profile()