""" data/database.py The base definitions for database access and definition """ from datetime import date, datetime from typing import Any, Sequence, TypeVar from sqlalchemy import ( create_engine, Date, func, Integer, MetaData, Row, text, ) from sqlalchemy.exc import IntegrityError, PendingRollbackError from sqlalchemy.orm import ( DeclarativeBase, Mapped, mapped_column, Session, sessionmaker, ) from sqlalchemy_utils import create_database, database_exists from util.environment import load_environment_variables project_variables = load_environment_variables() """ Databse naming conventions. ix == index uq == unique constraint ck == check constraint fk == foreign key pk == primary key """ DATABASE_NAMING_CONVENTION = { "ix": "ix_%(column_0_label)s", "uq": "uq_%(table_name)s_%(column_0_label)s", "ck": "ck_%(table_name)s_%(constraint_name)s", "fk": "fk_%(table_name)s_%(column_0_label)s_%(referred_table_name)s", "pk": "pk_%(table_name)s", } class Base(DeclarativeBase): metadata = MetaData(naming_convention=DATABASE_NAMING_CONVENTION) Table = TypeVar("Table", bound=Base) class PowerballDraw(Base): __tablename__ = 'powerball_draws' draw_date: Mapped[datetime] = mapped_column(Date, primary_key=True) main_ball1: Mapped[int] = mapped_column(Integer) main_ball2: Mapped[int] = mapped_column(Integer) main_ball3: Mapped[int] = mapped_column(Integer) main_ball4: Mapped[int] = mapped_column(Integer) main_ball5: Mapped[int] = mapped_column(Integer) powerball: Mapped[int] = mapped_column(Integer) power_play: Mapped[int] = mapped_column(Integer) def __repr__(self): return (f"PowerballDraw({self.draw_date=}, {self.main_ball1=}, " f"{self.main_ball2=}, {self.main_ball3=}, {self.main_ball4=}, " f"{self.main_ball5=}, {self.powerball=}, {self.power_play=})") def __str__(self): return (f"{self.draw_date}, [{self.main_ball1}, {self.main_ball2}, " f"{self.main_ball3}, {self.main_ball4}, {self.main_ball5}], " f"{self.powerball}, " f"{self.power_play if self.power_play else 1}x") def __eq__(self, other): return (self.draw_date == other.draw_date and self.main_ball1 == other.main_ball1 and \ self.main_ball2 == other.main_ball2 and self.main_ball3 == other.main_ball3 and \ self.main_ball4 == other.main_ball4 and self.main_ball5 == other.main_ball5 and \ self.powerball == other.powerball and self.power_play == other.power_play) class MegaMillionsDraw(Base): __tablename__ = 'mega_millions_draws' draw_date: Mapped[datetime] = mapped_column(Date, primary_key=True) main_ball1: Mapped[int] = mapped_column(Integer) main_ball2: Mapped[int] = mapped_column(Integer) main_ball3: Mapped[int] = mapped_column(Integer) main_ball4: Mapped[int] = mapped_column(Integer) main_ball5: Mapped[int] = mapped_column(Integer) mega_ball: Mapped[int] = mapped_column(Integer) megaplier: Mapped[int] = mapped_column(Integer) def __repr__(self): return (f"MegaMillionsDraw({self.draw_date=}, {self.main_ball1=}, " f"{self.main_ball2=}, {self.main_ball3=}, " f"{self.main_ball4=}, {self.main_ball5=}, " f"{self.mega_ball=}, {self.megaplier=})") def __str__(self): return (f"{self.draw_date}, [{self.main_ball1}, {self.main_ball2}, " f"{self.main_ball3}, {self.main_ball4}, " f"{self.main_ball5}], {self.mega_ball}, " f"{self.megaplier if self.megaplier else 1}x") def __eq__(self, other): return (self.draw_date == other.draw_date and self.main_ball1 == other.main_ball1 and \ self.main_ball2 == other.main_ball2 and self.main_ball3 == other.main_ball3 and \ self.main_ball4 == other.main_ball4 and self.main_ball5 == other.main_ball5 and \ self.mega_ball == other.mega_ball and self.megaplier == other.megaplier) class DataAccessLayer: def __init__(self): self.engine = None self.conn_string = project_variables["LP_DATABASE_URL"] def connect(self): self.engine = create_engine(self.conn_string) if not database_exists(self.engine.url): create_database(self.engine.url) Base.metadata.create_all(self.engine) self.Session = sessionmaker(bind=self.engine) class DataAccessLayer2: def __init__(self): self.engine = create_engine( project_variables["LP_DATABASE_URL"], echo=False, ) if not database_exists(self.engine.url): create_database(self.engine.url) Base.metadata.create_all(self.engine) self.session_local: sessionmaker[Session] = sessionmaker( bind=self.engine, ) self.session_factory = self.session_local def execute_query(self, query) -> Sequence[Row[Any]]: with self.session_factory() as session: try: result = session.execute(text(query)) return result.fetchall() finally: session.close() def add(self, record: Table) -> None: with self.session_factory() as session: try: session.add(record) session.commit() except IntegrityError: session.rollback() except PendingRollbackError: session.rollback() finally: session.close() def count(self, table_name: str, from_date: date | None = None) -> int: with self.session_factory() as session: if from_date: return session.query(table_name).where( table_name.draw_date >= from_date, ).count() else: return session.query(table_name).count() def most_recent(self, table_name: Table) -> Row[Table] | None: with self.session_factory() as session: return session.query(func.max(table_name.draw_date)).first() dal = DataAccessLayer() dal2 = DataAccessLayer2()