Add some helper types for CRUD infra
This commit is contained in:
parent
1c08151558
commit
300cb7bd67
1 changed files with 82 additions and 2 deletions
|
|
@ -1,6 +1,8 @@
|
|||
from sqlalchemy.orm import DeclarativeBase
|
||||
from typing import Callable
|
||||
from sqlalchemy.orm import DeclarativeBase, Session
|
||||
from src.config.database import Database as DatabaseConfig
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy import create_engine, Engine
|
||||
import abc
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
|
|
@ -13,3 +15,81 @@ def get_database(config: DatabaseConfig):
|
|||
Base.metadata.create_all(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()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue