Restructure where the CRUD base class lives for clarity and to avoid circular deps
This commit is contained in:
parent
0d884e5836
commit
c3cd7c2796
8 changed files with 40 additions and 37 deletions
|
|
@ -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',
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
20
src/services/repo.py
Normal file
20
src/services/repo.py
Normal 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
|
||||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from src.infra.db import CRUD
|
||||
from src.services.repo import CRUD
|
||||
import abc
|
||||
from src.config.email import EmailAddress
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue