""" 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, engine, func, Integer, MetaData, Row, text from sqlalchemy.exc import IntegrityError, PendingRollbackError from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, sessionmaker, Session 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=}, {self.main_ball2=}, {self.main_ball3=}, " \ f"{self.main_ball4=}, {self.main_ball5=}, {self.powerball=}, {self.power_play=})" def __str__(self): return f"{self.draw_date}, [{self.main_ball1}, {self.main_ball2}, {self.main_ball3}, {self.main_ball4}, " \ f"{self.main_ball5}], {self.powerball}, {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=}, {self.main_ball2=}, {self.main_ball3=}, " \ f"{self.main_ball4=}, {self.main_ball5=}, {self.mega_ball=}, {self.megaplier=})" def __str__(self): return f"{self.draw_date}, [{self.main_ball1}, {self.main_ball2}, {self.main_ball3}, {self.main_ball4}, " \ f"{self.main_ball5}], {self.mega_ball}, {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()