Compare commits

..

No commits in common. "b601a8fbe85fc25f12ea6c2b2f0e6dd34264b0a5" and "1c081515584d4f00af85ca53401f17725d15b3d0" have entirely different histories.

6 changed files with 106 additions and 145 deletions

View file

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

View file

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

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, 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:

View file

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

View file

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

View file

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