Restructure where the CRUD base class lives for clarity and to avoid circular deps

This commit is contained in:
Campbell Alden 2026-09-05 23:38:46 +09:00
parent 0d884e5836
commit c3cd7c2796
8 changed files with 40 additions and 37 deletions

View file

@ -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 .users import UserModel, UserRepoImpl
from .publication import ( from .publication import (
PublicationEntryModel, PublicationEntryModel,
@ -11,7 +11,6 @@ from .subscription import SubscriptionModel, SubscriptionRepoImpl
__all__ = [ __all__ = [
'Base', 'Base',
'CRUD',
'CRUDRepo', 'CRUDRepo',
'UpdateMissingEntryError', 'UpdateMissingEntryError',
'get_database', 'get_database',

View file

@ -1,8 +1,8 @@
from src.services.repo import CRUD
from typing import Callable from typing import Callable
from sqlalchemy.orm import DeclarativeBase, Session 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, Engine
import abc
class Base(DeclarativeBase): class Base(DeclarativeBase):
@ -22,25 +22,7 @@ class UpdateMissingEntryError(Exception):
super().__init__(f'Attempted to update non-existing {model_name} {id}') super().__init__(f'Attempted to update non-existing {model_name} {id}')
class CRUD[T, C](abc.ABC): class CRUDRepoImpl[T: Base, U, C](CRUD[U, C]):
@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__( def __init__(
self, self,
db: Engine, db: Engine,

View file

@ -1,12 +1,12 @@
from sqlalchemy import Engine from sqlalchemy import Engine
from typing import Callable from typing import Callable
from src.infra import CRUDRepo, Base from src.infra import CRUDRepoImpl, Base
import pytest import pytest
class CRUDRepoHelper[A: Base, B, C]: class CRUDRepoHelper[A: Base, B, C]:
@pytest.fixture @pytest.fixture
def repo(self, db: Engine) -> CRUDRepo[A, B, C]: def repo(self, db: Engine) -> CRUDRepoImpl[A, B, C]:
import pdb import pdb
pdb.set_trace() pdb.set_trace()
@ -17,7 +17,7 @@ class CRUDRepoHelper[A: Base, B, C]:
raise NotImplementedError() raise NotImplementedError()
@pytest.fixture @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) return repo.create(dto, None)
@pytest.fixture @pytest.fixture
@ -31,26 +31,28 @@ class CRUDRepoHelper[A: Base, B, C]:
def test_create(self, item: B): def test_create(self, item: B):
assert item is not None 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): def always_invalid(*_args, **_kw):
raise ValueError('test') raise ValueError('test')
with pytest.raises(ValueError): with pytest.raises(ValueError):
repo.create(dto, always_invalid) 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 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 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 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 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 assert repo.get(1) is not None
repo.delete(1, None) repo.delete(1, None)
assert repo.get(1) is None assert repo.get(1) is None

View file

@ -5,7 +5,7 @@ from typing import TYPE_CHECKING
from datetime import datetime, timedelta from datetime import datetime, timedelta
from src.infra.users import UserModel from src.infra.users import UserModel
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, CRUDRepoImpl
from sqlalchemy import Engine, ForeignKey, CheckConstraint, DateTime, func, UniqueConstraint, select from sqlalchemy import Engine, ForeignKey, CheckConstraint, DateTime, func, UniqueConstraint, select
from src.services.subscription import ( from src.services.subscription import (
SubscriptionRepo, SubscriptionRepo,
@ -62,7 +62,7 @@ def mutate_subscription(db_sub: SubscriptionModel, sub: Subscription):
db_sub.sequence_notified = sub.sequence_notified 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): def __init__(self, db: Engine):
super().__init__( super().__init__(
db, SubscriptionModel, model_to_subscription, mutate_subscription, lambda u: SubscriptionModel(**asdict(u)) db, SubscriptionModel, model_to_subscription, mutate_subscription, lambda u: SubscriptionModel(**asdict(u))

View file

@ -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, CRUDRepoImpl
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
@ -52,7 +52,7 @@ def make_create_user(hasher: PasswordHasher):
return create_user return create_user
class UserRepoImpl(CRUDRepo[UserModel, User, UserDTO], UserRepo): class UserRepoImpl(CRUDRepoImpl[UserModel, User, UserDTO], UserRepo):
def __init__(self, db: Engine): def __init__(self, db: Engine):
hasher = PasswordHasher() hasher = PasswordHasher()
super().__init__(db, UserModel, model_to_user, mutate_user, make_create_user(hasher)) super().__init__(db, UserModel, model_to_user, mutate_user, make_create_user(hasher))

20
src/services/repo.py Normal file
View file

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

View file

@ -1,4 +1,4 @@
from src.infra.db import CRUD from src.services.repo import CRUD
import abc import abc
from datetime import datetime from datetime import datetime
from src.services.publications import ( from src.services.publications import (

View file

@ -1,4 +1,4 @@
from src.infra.db import CRUD from src.services.repo import CRUD
import abc import abc
from src.config.email import EmailAddress from src.config.email import EmailAddress