Compare commits

...

3 commits

Author SHA1 Message Date
Campbell Alden
b601a8fbe8 Flesh out more of the subscription infra and services using CRUD types 2026-08-24 00:12:47 +09:00
Campbell Alden
bdf74c2dbd Update users to use the new CRUD types 2026-08-24 00:12:29 +09:00
Campbell Alden
300cb7bd67 Add some helper types for CRUD infra 2026-08-24 00:12:17 +09:00
6 changed files with 145 additions and 106 deletions

View file

@ -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()

View file

@ -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 {}

View file

@ -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,42 +33,7 @@ def model_to_user(db_user: UserModel) -> User:
)
class UserRepoImpl(UserRepo):
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
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()
@ -78,7 +43,28 @@ class UserRepoImpl(UserRepo):
db_user.email = user.email
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:
full_user = self.get_user_by_email(user.email)

View file

@ -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)

View file

@ -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

View file

@ -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
def update(user: User) -> User:
user.email_confirmed = True
self._repo.update_user(user)
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()