Files
Lottery_Project/data/database.py
T
2026-05-27 08:59:05 -04:00

192 lines
6.2 KiB
Python

"""
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()