Compare commits
3 commits
1c08151558
...
b601a8fbe8
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b601a8fbe8 | ||
|
|
bdf74c2dbd | ||
|
|
300cb7bd67 |
6 changed files with 145 additions and 106 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue