192 lines
6.2 KiB
Python
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()
|