From b601a8fbe85fc25f12ea6c2b2f0e6dd34264b0a5 Mon Sep 17 00:00:00 2001 From: Campbell Alden Date: Mon, 24 Aug 2026 00:12:47 +0900 Subject: [PATCH] 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)