diff --git a/src/infra/db.py b/src/infra/db.py index 841cbb8..6359df4 100644 --- a/src/infra/db.py +++ b/src/infra/db.py @@ -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()