From c3cd7c279613b63d742d96194eb175baac9f7833 Mon Sep 17 00:00:00 2001 From: Campbell Alden Date: Sat, 5 Sep 2026 23:38:46 +0900 Subject: [PATCH] Restructure where the CRUD base class lives for clarity and to avoid circular deps --- src/infra/__init__.py | 3 +-- src/infra/db.py | 22 ++-------------------- src/infra/db_tests.py | 20 +++++++++++--------- src/infra/subscription.py | 4 ++-- src/infra/users.py | 4 ++-- src/services/repo.py | 20 ++++++++++++++++++++ src/services/subscription.py | 2 +- src/services/users/repo.py | 2 +- 8 files changed, 40 insertions(+), 37 deletions(-) create mode 100644 src/services/repo.py diff --git a/src/infra/__init__.py b/src/infra/__init__.py index 8851384..ee15f85 100644 --- a/src/infra/__init__.py +++ b/src/infra/__init__.py @@ -1,4 +1,4 @@ -from .db import Base, CRUD, CRUDRepo, UpdateMissingEntryError, get_database +from .db import Base, CRUDRepoImpl, UpdateMissingEntryError, get_database from .users import UserModel, UserRepoImpl from .publication import ( PublicationEntryModel, @@ -11,7 +11,6 @@ from .subscription import SubscriptionModel, SubscriptionRepoImpl __all__ = [ 'Base', - 'CRUD', 'CRUDRepo', 'UpdateMissingEntryError', 'get_database', diff --git a/src/infra/db.py b/src/infra/db.py index 6359df4..ad9ff98 100644 --- a/src/infra/db.py +++ b/src/infra/db.py @@ -1,8 +1,8 @@ +from src.services.repo import CRUD from typing import Callable from sqlalchemy.orm import DeclarativeBase, Session from src.config.database import Database as DatabaseConfig from sqlalchemy import create_engine, Engine -import abc class Base(DeclarativeBase): @@ -22,25 +22,7 @@ class UpdateMissingEntryError(Exception): 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]): +class CRUDRepoImpl[T: Base, U, C](CRUD[U, C]): def __init__( self, db: Engine, diff --git a/src/infra/db_tests.py b/src/infra/db_tests.py index e764924..01ecf4f 100644 --- a/src/infra/db_tests.py +++ b/src/infra/db_tests.py @@ -1,12 +1,12 @@ from sqlalchemy import Engine from typing import Callable -from src.infra import CRUDRepo, Base +from src.infra import CRUDRepoImpl, Base import pytest class CRUDRepoHelper[A: Base, B, C]: @pytest.fixture - def repo(self, db: Engine) -> CRUDRepo[A, B, C]: + def repo(self, db: Engine) -> CRUDRepoImpl[A, B, C]: import pdb pdb.set_trace() @@ -17,7 +17,7 @@ class CRUDRepoHelper[A: Base, B, C]: raise NotImplementedError() @pytest.fixture - def item(self, repo: CRUDRepo[A, B, C], dto: C) -> B: + def item(self, repo: CRUDRepoImpl[A, B, C], dto: C) -> B: return repo.create(dto, None) @pytest.fixture @@ -31,26 +31,28 @@ class CRUDRepoHelper[A: Base, B, C]: def test_create(self, item: B): assert item is not None - def test_create_returns_none_if_invalid(self, repo: CRUDRepo[A, B, C], dto: C): + def test_create_returns_none_if_invalid(self, repo: CRUDRepoImpl[A, B, C], dto: C): def always_invalid(*_args, **_kw): raise ValueError('test') with pytest.raises(ValueError): repo.create(dto, always_invalid) - def test_create_returns_a_value_if_valid(self, repo: CRUDRepo[A, B, C], dto: C): + def test_create_returns_a_value_if_valid(self, repo: CRUDRepoImpl[A, B, C], dto: C): assert repo.create(dto, lambda _: None) is not None - def test_get_returns_none_when_missing(self, repo: CRUDRepo[A, B, C]): + def test_get_returns_none_when_missing(self, repo: CRUDRepoImpl[A, B, C]): assert repo.get(42) is None - def test_get_returns_a_value_if_it_exists(self, repo: CRUDRepo[A, B, C], item: B): + def test_get_returns_a_value_if_it_exists(self, repo: CRUDRepoImpl[A, B, C], item: B): assert repo.get(1) == item - def test_update(self, repo: CRUDRepo[A, B, C], transform: Callable[[B], B], item: B, expected_transformed_value: B): + def test_update( + self, repo: CRUDRepoImpl[A, B, C], transform: Callable[[B], B], item: B, expected_transformed_value: B + ): assert repo.update(1, transform) == expected_transformed_value - def test_delete(self, repo: CRUDRepo[A, B, C], item: B): + def test_delete(self, repo: CRUDRepoImpl[A, B, C], item: B): assert repo.get(1) is not None repo.delete(1, None) assert repo.get(1) is None diff --git a/src/infra/subscription.py b/src/infra/subscription.py index 6c8131d..c696be3 100644 --- a/src/infra/subscription.py +++ b/src/infra/subscription.py @@ -5,7 +5,7 @@ from typing import TYPE_CHECKING from datetime import datetime, timedelta from src.infra.users import UserModel from sqlalchemy.orm import Mapped, mapped_column, relationship, Session, joinedload -from src.infra.db import Base, CRUDRepo +from src.infra.db import Base, CRUDRepoImpl from sqlalchemy import Engine, ForeignKey, CheckConstraint, DateTime, func, UniqueConstraint, select from src.services.subscription import ( SubscriptionRepo, @@ -62,7 +62,7 @@ def mutate_subscription(db_sub: SubscriptionModel, sub: Subscription): db_sub.sequence_notified = sub.sequence_notified -class SubscriptionRepoImpl(SubscriptionRepo, CRUDRepo[SubscriptionModel, Subscription, SubscriptionCreateParams]): +class SubscriptionRepoImpl(SubscriptionRepo, CRUDRepoImpl[SubscriptionModel, Subscription, SubscriptionCreateParams]): def __init__(self, db: Engine): super().__init__( db, SubscriptionModel, model_to_subscription, mutate_subscription, lambda u: SubscriptionModel(**asdict(u)) diff --git a/src/infra/users.py b/src/infra/users.py index aa17d43..a4c7a14 100644 --- a/src/infra/users.py +++ b/src/infra/users.py @@ -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, CRUDRepoImpl from src.services.users.data import UserDTO, User from src.services.users.repo import UserRepo @@ -52,7 +52,7 @@ def make_create_user(hasher: PasswordHasher): return create_user -class UserRepoImpl(CRUDRepo[UserModel, User, UserDTO], UserRepo): +class UserRepoImpl(CRUDRepoImpl[UserModel, User, UserDTO], UserRepo): def __init__(self, db: Engine): hasher = PasswordHasher() super().__init__(db, UserModel, model_to_user, mutate_user, make_create_user(hasher)) diff --git a/src/services/repo.py b/src/services/repo.py new file mode 100644 index 0000000..dceec86 --- /dev/null +++ b/src/services/repo.py @@ -0,0 +1,20 @@ +from typing import Callable +import abc + + +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 diff --git a/src/services/subscription.py b/src/services/subscription.py index 856ad6c..720f747 100644 --- a/src/services/subscription.py +++ b/src/services/subscription.py @@ -1,4 +1,4 @@ -from src.infra.db import CRUD +from src.services.repo import CRUD import abc from datetime import datetime from src.services.publications import ( diff --git a/src/services/users/repo.py b/src/services/users/repo.py index 6f8f40c..cd6bc74 100644 --- a/src/services/users/repo.py +++ b/src/services/users/repo.py @@ -1,4 +1,4 @@ -from src.infra.db import CRUD +from src.services.repo import CRUD import abc from src.config.email import EmailAddress