From 300cb7bd67c8fc1cd7a20c37ff3876a747ad70b5 Mon Sep 17 00:00:00 2001 From: Campbell Alden Date: Mon, 24 Aug 2026 00:12:17 +0900 Subject: [PATCH 1/3] Add some helper types for CRUD infra --- src/infra/db.py | 84 +++++++++++++++++++++++++++++++++++++++++++++++-- 1 file changed, 82 insertions(+), 2 deletions(-) diff --git a/src/infra/db.py b/src/infra/db.py index 841cbb8..6359df4 100644 --- a/src/infra/db.py +++ b/src/infra/db.py @@ -1,6 +1,8 @@ -from sqlalchemy.orm import DeclarativeBase +from typing import Callable +from sqlalchemy.orm import DeclarativeBase, Session from src.config.database import Database as DatabaseConfig -from sqlalchemy import create_engine +from sqlalchemy import create_engine, Engine +import abc class Base(DeclarativeBase): @@ -13,3 +15,81 @@ 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() From bdf74c2dbd22d969c2044ad36c6922beec2e3d0b Mon Sep 17 00:00:00 2001 From: Campbell Alden Date: Mon, 24 Aug 2026 00:12:29 +0900 Subject: [PATCH 2/3] Update users to use the new CRUD types --- src/infra/users.py | 62 ++++++++++++++--------------------- src/services/users/repo.py | 15 ++------- src/services/users/service.py | 16 ++++++--- 3 files changed, 37 insertions(+), 56 deletions(-) diff --git a/src/infra/users.py b/src/infra/users.py index f358acb..a01e4bc 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 +from src.infra.db import Base, CRUDRepo from src.services.users.data import UserDTO, User from src.services.users.repo import UserRepo @@ -33,27 +33,31 @@ def model_to_user(db_user: UserModel) -> User: ) -class UserRepoImpl(UserRepo): +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): def __init__(self, db: Engine): - self._db = db - self._hasher = PasswordHasher() + hasher = PasswordHasher() + super().__init__(db, UserModel, model_to_user, mutate_user, make_create_user(hasher)) 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) + self._hasher = hasher def get_user_by_email(self, email: EmailAddress) -> User | None: with Session(self._db) as session: @@ -62,24 +66,6 @@ class UserRepoImpl(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/users/repo.py b/src/services/users/repo.py index f50836b..6f8f40c 100644 --- a/src/services/users/repo.py +++ b/src/services/users/repo.py @@ -1,26 +1,15 @@ +from src.infra.db import CRUD import abc from src.config.email import EmailAddress from .data import 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 - +class UserRepo(CRUD[User, UserDTO]): @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 d481cc0..ca916b6 100644 --- a/src/services/users/service.py +++ b/src/services/users/service.py @@ -65,14 +65,20 @@ class UserService: if not claim or claim.sub != user_id: 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 # the token invalid for this request. if not user or user.email != claim.email or user.email_confirmed: raise EmailConfirmationTokenInvalid - user.email_confirmed = True - self._repo.update_user(user) + def update(user: User) -> User: + user.email_confirmed = True + return user + + self._repo.update( + user_id, + update, + ) def _send_confirmation_email(self, user: User): env = Environment(loader=PackageLoader('src'), autoescape=select_autoescape()) @@ -96,13 +102,13 @@ class UserService: raise SignupError('The given password was not acceptable') # Create a user in persistence - created_user = self._repo.create_user(user) + created_user = self._repo.create(user, None) # 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_by_id(user_id) + user = self._repo.get(user_id) if user: return user.to_profile() From b601a8fbe85fc25f12ea6c2b2f0e6dd34264b0a5 Mon Sep 17 00:00:00 2001 From: Campbell Alden Date: Mon, 24 Aug 2026 00:12:47 +0900 Subject: [PATCH 3/3] Flesh out more of the subscription infra and services using CRUD types --- src/infra/subscription.py | 36 +++++++++++++++++----------------- src/services/subscription.py | 38 ++++++++---------------------------- 2 files changed, 26 insertions(+), 48 deletions(-) diff --git a/src/infra/subscription.py b/src/infra/subscription.py index 40da20f..5575100 100644 --- a/src/infra/subscription.py +++ b/src/infra/subscription.py @@ -1,15 +1,17 @@ +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, model_to_user +from src.infra.users import UserModel from sqlalchemy.orm import Mapped, mapped_column, relationship, Session, joinedload -from src.infra.db import Base -from sqlalchemy import Engine, ForeignKey, CheckConstraint, DateTime, func, UniqueConstraint, select +from src.infra.db import Base, CRUDRepo +from sqlalchemy import Engine, ForeignKey, CheckConstraint, DateTime, func, UniqueConstraint, select, delete from src.services.subscription import ( SubscriptionRepo, Subscription, SubscriptionId, - UpdateSubscription, + SubscriptionCreateParams, AvailableSubscriptionEntry, ) @@ -46,16 +48,24 @@ class SubscriptionModel(Base): def model_to_subscription(db_sub: SubscriptionModel) -> Subscription: return Subscription( id=db_sub.id, - user=model_to_user(db_sub.user).to_profile(), + user_id=db_sub.user_id, publication_order=model_to_order(db_sub.order).to_profile(), sequence_seen=db_sub.sequence_seen, start=db_sub.start, ) -class SubscriptionRepoImpl(SubscriptionRepo): +def mutate_subscription(db_sub: SubscriptionModel, sub: Subscription): + pass + + +class SubscriptionRepoImpl(SubscriptionRepo, CRUDRepo[SubscriptionModel, Subscription, SubscriptionCreateParams]): def __init__(self, db: Engine): - self._db = db + super().__init__( + db, SubscriptionModel, model_to_subscription, mutate_subscription, lambda u: SubscriptionModel(**asdict(u)) + ) + + self._logger = logging.getLogger('UserRepoImpl') def get_subscriptions_for_user(self, user_id: int) -> list[Subscription]: with Session(self._db) as session: @@ -65,7 +75,6 @@ class SubscriptionRepoImpl(SubscriptionRepo): .where(SubscriptionModel.user_id == user_id) .options( joinedload(SubscriptionModel.order), - joinedload(SubscriptionModel.user), ) ) .unique() @@ -78,21 +87,12 @@ class SubscriptionRepoImpl(SubscriptionRepo): db_sub = session.get( SubscriptionModel, sub_id, - options=[joinedload(SubscriptionModel.order), joinedload(SubscriptionModel.user)], + options=[joinedload(SubscriptionModel.order)], ) 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/services/subscription.py b/src/services/subscription.py index 25f70a7..5eb5d51 100644 --- a/src/services/subscription.py +++ b/src/services/subscription.py @@ -1,3 +1,4 @@ +from src.infra.db import CRUD import abc from datetime import datetime from src.services.publications import ( @@ -5,7 +6,6 @@ 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: UserProfile + user_id: int publication_order: PublicationOrderProfile sequence_seen: int start: datetime @@ -29,8 +29,6 @@ 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 @@ -43,27 +41,11 @@ class AvailableSubscriptionEntry: sequence: PublicationSequence -class SubscriptionRepo(abc.ABC): +class SubscriptionRepo(CRUD[Subscription, SubscriptionCreateParams]): @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 @@ -91,15 +73,11 @@ class SubscriptionService: return self._repo.get_available_entries(user_id) def create_suscription(self, params: SubscriptionCreateParams) -> Subscription: - return self._repo.create_subscription_for_user(params.user_id, params.order_id) + return self._repo.create(params, None) def delete_subcription(self, user_id: int, subscription_id: SubscriptionId): - 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') + def verify(subscription: Subscription): + if subscription.user_id != user_id: + raise ValueError('Only the owning user can delete a subscription') - self._repo.delete_subscription(subscription_id) + self._repo.delete(subscription_id, verify)