Add some helper types for CRUD infra

This commit is contained in:
Campbell Alden 2026-08-24 00:12:17 +09:00
parent 1c08151558
commit 300cb7bd67

View file

@ -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 src.config.database import Database as DatabaseConfig
from sqlalchemy import create_engine from sqlalchemy import create_engine, Engine
import abc
class Base(DeclarativeBase): class Base(DeclarativeBase):
@ -13,3 +15,81 @@ def get_database(config: DatabaseConfig):
Base.metadata.create_all(engine) Base.metadata.create_all(engine)
return 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()