Compare commits
No commits in common. "b601a8fbe85fc25f12ea6c2b2f0e6dd34264b0a5" and "1c081515584d4f00af85ca53401f17725d15b3d0" have entirely different histories.
b601a8fbe8
...
1c08151558
6 changed files with 106 additions and 145 deletions
|
|
@ -1,8 +1,6 @@
|
||||||
from typing import Callable
|
from sqlalchemy.orm import DeclarativeBase
|
||||||
from sqlalchemy.orm import DeclarativeBase, Session
|
|
||||||
from src.config.database import Database as DatabaseConfig
|
from src.config.database import Database as DatabaseConfig
|
||||||
from sqlalchemy import create_engine, Engine
|
from sqlalchemy import create_engine
|
||||||
import abc
|
|
||||||
|
|
||||||
|
|
||||||
class Base(DeclarativeBase):
|
class Base(DeclarativeBase):
|
||||||
|
|
@ -15,81 +13,3 @@ def get_database(config: DatabaseConfig):
|
||||||
Base.metadata.create_all(engine)
|
Base.metadata.create_all(engine)
|
||||||
|
|
||||||
return 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,17 +1,15 @@
|
||||||
from dataclasses import asdict
|
|
||||||
import logging
|
|
||||||
from src.infra.publication import model_to_order
|
from src.infra.publication import model_to_order
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
from datetime import datetime
|
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 sqlalchemy.orm import Mapped, mapped_column, relationship, Session, joinedload
|
||||||
from src.infra.db import Base, CRUDRepo
|
from src.infra.db import Base
|
||||||
from sqlalchemy import Engine, ForeignKey, CheckConstraint, DateTime, func, UniqueConstraint, select, delete
|
from sqlalchemy import Engine, ForeignKey, CheckConstraint, DateTime, func, UniqueConstraint, select
|
||||||
from src.services.subscription import (
|
from src.services.subscription import (
|
||||||
SubscriptionRepo,
|
SubscriptionRepo,
|
||||||
Subscription,
|
Subscription,
|
||||||
SubscriptionId,
|
SubscriptionId,
|
||||||
SubscriptionCreateParams,
|
UpdateSubscription,
|
||||||
AvailableSubscriptionEntry,
|
AvailableSubscriptionEntry,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -48,24 +46,16 @@ class SubscriptionModel(Base):
|
||||||
def model_to_subscription(db_sub: SubscriptionModel) -> Subscription:
|
def model_to_subscription(db_sub: SubscriptionModel) -> Subscription:
|
||||||
return Subscription(
|
return Subscription(
|
||||||
id=db_sub.id,
|
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(),
|
publication_order=model_to_order(db_sub.order).to_profile(),
|
||||||
sequence_seen=db_sub.sequence_seen,
|
sequence_seen=db_sub.sequence_seen,
|
||||||
start=db_sub.start,
|
start=db_sub.start,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def mutate_subscription(db_sub: SubscriptionModel, sub: Subscription):
|
class SubscriptionRepoImpl(SubscriptionRepo):
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
class SubscriptionRepoImpl(SubscriptionRepo, CRUDRepo[SubscriptionModel, Subscription, SubscriptionCreateParams]):
|
|
||||||
def __init__(self, db: Engine):
|
def __init__(self, db: Engine):
|
||||||
super().__init__(
|
self._db = db
|
||||||
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]:
|
def get_subscriptions_for_user(self, user_id: int) -> list[Subscription]:
|
||||||
with Session(self._db) as session:
|
with Session(self._db) as session:
|
||||||
|
|
@ -75,6 +65,7 @@ class SubscriptionRepoImpl(SubscriptionRepo, CRUDRepo[SubscriptionModel, Subscri
|
||||||
.where(SubscriptionModel.user_id == user_id)
|
.where(SubscriptionModel.user_id == user_id)
|
||||||
.options(
|
.options(
|
||||||
joinedload(SubscriptionModel.order),
|
joinedload(SubscriptionModel.order),
|
||||||
|
joinedload(SubscriptionModel.user),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
.unique()
|
.unique()
|
||||||
|
|
@ -87,12 +78,21 @@ class SubscriptionRepoImpl(SubscriptionRepo, CRUDRepo[SubscriptionModel, Subscri
|
||||||
db_sub = session.get(
|
db_sub = session.get(
|
||||||
SubscriptionModel,
|
SubscriptionModel,
|
||||||
sub_id,
|
sub_id,
|
||||||
options=[joinedload(SubscriptionModel.order)],
|
options=[joinedload(SubscriptionModel.order), joinedload(SubscriptionModel.user)],
|
||||||
)
|
)
|
||||||
|
|
||||||
if db_sub:
|
if db_sub:
|
||||||
return model_to_subscription(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]]:
|
def get_available_entries(self, user_id: int) -> dict[SubscriptionId, list[AvailableSubscriptionEntry]]:
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -6,7 +6,7 @@ from src.config.email import EmailAddress
|
||||||
from sqlalchemy import String, VARCHAR, Engine, select
|
from sqlalchemy import String, VARCHAR, Engine, select
|
||||||
from sqlalchemy.orm import mapped_column, Mapped, Session, relationship
|
from sqlalchemy.orm import mapped_column, Mapped, Session, relationship
|
||||||
from argon2 import PasswordHasher
|
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.data import UserDTO, User
|
||||||
from src.services.users.repo import UserRepo
|
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):
|
class UserRepoImpl(UserRepo):
|
||||||
# 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):
|
def __init__(self, db: Engine):
|
||||||
hasher = PasswordHasher()
|
self._db = db
|
||||||
super().__init__(db, UserModel, model_to_user, mutate_user, make_create_user(hasher))
|
self._hasher = PasswordHasher()
|
||||||
self._logger = logging.getLogger('UserRepoImpl')
|
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:
|
def get_user_by_email(self, email: EmailAddress) -> User | None:
|
||||||
with Session(self._db) as session:
|
with Session(self._db) as session:
|
||||||
|
|
@ -66,6 +62,24 @@ class UserRepoImpl(CRUDRepo[UserModel, User, UserDTO], UserRepo):
|
||||||
if user:
|
if user:
|
||||||
return model_to_user(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:
|
def auth_as_user(self, user: UserDTO) -> User | None:
|
||||||
full_user = self.get_user_by_email(user.email)
|
full_user = self.get_user_by_email(user.email)
|
||||||
if full_user is None:
|
if full_user is None:
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,3 @@
|
||||||
from src.infra.db import CRUD
|
|
||||||
import abc
|
import abc
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from src.services.publications import (
|
from src.services.publications import (
|
||||||
|
|
@ -6,6 +5,7 @@ from src.services.publications import (
|
||||||
Sequence as PublicationSequence,
|
Sequence as PublicationSequence,
|
||||||
Publication,
|
Publication,
|
||||||
)
|
)
|
||||||
|
from src.services.users.data import UserProfile
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
|
||||||
SubscriptionId = int
|
SubscriptionId = int
|
||||||
|
|
@ -20,7 +20,7 @@ class SubscriptionCreateParams:
|
||||||
@dataclass
|
@dataclass
|
||||||
class Subscription:
|
class Subscription:
|
||||||
id: SubscriptionId
|
id: SubscriptionId
|
||||||
user_id: int
|
user: UserProfile
|
||||||
publication_order: PublicationOrderProfile
|
publication_order: PublicationOrderProfile
|
||||||
sequence_seen: int
|
sequence_seen: int
|
||||||
start: datetime
|
start: datetime
|
||||||
|
|
@ -29,6 +29,8 @@ class Subscription:
|
||||||
@dataclass
|
@dataclass
|
||||||
class UpdateSubscription:
|
class UpdateSubscription:
|
||||||
id: int
|
id: int
|
||||||
|
user_id: int
|
||||||
|
publication_order_id: int | None = None
|
||||||
sequence_seen: int | None = None
|
sequence_seen: int | None = None
|
||||||
start: datetime | None = None
|
start: datetime | None = None
|
||||||
|
|
||||||
|
|
@ -41,11 +43,27 @@ class AvailableSubscriptionEntry:
|
||||||
sequence: PublicationSequence
|
sequence: PublicationSequence
|
||||||
|
|
||||||
|
|
||||||
class SubscriptionRepo(CRUD[Subscription, SubscriptionCreateParams]):
|
class SubscriptionRepo(abc.ABC):
|
||||||
@abc.abstractmethod
|
@abc.abstractmethod
|
||||||
def get_subscriptions_for_user(self, user_id: int) -> list[Subscription]:
|
def get_subscriptions_for_user(self, user_id: int) -> list[Subscription]:
|
||||||
pass
|
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
|
@abc.abstractmethod
|
||||||
def get_available_entries(self, user_id: int) -> dict[SubscriptionId, list[AvailableSubscriptionEntry]]:
|
def get_available_entries(self, user_id: int) -> dict[SubscriptionId, list[AvailableSubscriptionEntry]]:
|
||||||
pass
|
pass
|
||||||
|
|
@ -73,11 +91,15 @@ class SubscriptionService:
|
||||||
return self._repo.get_available_entries(user_id)
|
return self._repo.get_available_entries(user_id)
|
||||||
|
|
||||||
def create_suscription(self, params: SubscriptionCreateParams) -> Subscription:
|
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 delete_subcription(self, user_id: int, subscription_id: SubscriptionId):
|
||||||
def verify(subscription: Subscription):
|
subscription = self._repo.get_subscription_by_id(subscription_id)
|
||||||
if subscription.user_id != user_id:
|
if subscription is None:
|
||||||
raise ValueError('Only the owning user can delete a subscription')
|
# 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)
|
||||||
|
|
|
||||||
|
|
@ -1,15 +1,26 @@
|
||||||
from src.infra.db import CRUD
|
|
||||||
import abc
|
import abc
|
||||||
from src.config.email import EmailAddress
|
from src.config.email import EmailAddress
|
||||||
|
|
||||||
from .data import User, UserDTO
|
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
|
@abc.abstractmethod
|
||||||
def get_user_by_email(self, email: EmailAddress) -> User | None:
|
def get_user_by_email(self, email: EmailAddress) -> User | None:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
@abc.abstractmethod
|
||||||
|
def update_user(self, user: User):
|
||||||
|
pass
|
||||||
|
|
||||||
@abc.abstractmethod
|
@abc.abstractmethod
|
||||||
def auth_as_user(self, user: UserDTO) -> User | None:
|
def auth_as_user(self, user: UserDTO) -> User | None:
|
||||||
pass
|
pass
|
||||||
|
|
|
||||||
|
|
@ -65,20 +65,14 @@ class UserService:
|
||||||
if not claim or claim.sub != user_id:
|
if not claim or claim.sub != user_id:
|
||||||
raise EmailConfirmationTokenInvalid
|
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
|
# 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.
|
# the token invalid for this request.
|
||||||
if not user or user.email != claim.email or user.email_confirmed:
|
if not user or user.email != claim.email or user.email_confirmed:
|
||||||
raise EmailConfirmationTokenInvalid
|
raise EmailConfirmationTokenInvalid
|
||||||
|
|
||||||
def update(user: User) -> User:
|
user.email_confirmed = True
|
||||||
user.email_confirmed = True
|
self._repo.update_user(user)
|
||||||
return user
|
|
||||||
|
|
||||||
self._repo.update(
|
|
||||||
user_id,
|
|
||||||
update,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _send_confirmation_email(self, user: User):
|
def _send_confirmation_email(self, user: User):
|
||||||
env = Environment(loader=PackageLoader('src'), autoescape=select_autoescape())
|
env = Environment(loader=PackageLoader('src'), autoescape=select_autoescape())
|
||||||
|
|
@ -102,13 +96,13 @@ class UserService:
|
||||||
raise SignupError('The given password was not acceptable')
|
raise SignupError('The given password was not acceptable')
|
||||||
|
|
||||||
# Create a user in persistence
|
# Create a user in persistence
|
||||||
created_user = self._repo.create(user, None)
|
created_user = self._repo.create_user(user)
|
||||||
# send a confirmation email
|
# send a confirmation email
|
||||||
self._send_confirmation_email(created_user)
|
self._send_confirmation_email(created_user)
|
||||||
|
|
||||||
return created_user.to_profile()
|
return created_user.to_profile()
|
||||||
|
|
||||||
def get_user_by_id(self, user_id: int) -> UserProfile | None:
|
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:
|
if user:
|
||||||
return user.to_profile()
|
return user.to_profile()
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue