From 3fe2e1bcd5500b7de01cef2d9606a81802f16698 Mon Sep 17 00:00:00 2001 From: Chris Smith Date: Tue, 26 May 2026 19:19:44 -0400 Subject: [PATCH] Add new DAL and update code to create missing database --- data/database.py | 85 ++++++++++++++++++++++++++++++++++++++---------- 1 file changed, 68 insertions(+), 17 deletions(-) diff --git a/data/database.py b/data/database.py index 4b84eaf..5bd1a3b 100644 --- a/data/database.py +++ b/data/database.py @@ -3,10 +3,13 @@ The base definitions for database access and definition """ -from datetime import datetime +from datetime import date, datetime +from typing import Any, Sequence, TypeVar -from sqlalchemy import create_engine, Date, Integer, MetaData, String -from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, sessionmaker +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 @@ -32,20 +35,7 @@ class Base(DeclarativeBase): metadata = MetaData(naming_convention=DATABASE_NAMING_CONVENTION) -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) - Base.metadata.create_all(self.engine) - self.Session = sessionmaker(bind=self.engine) - - -dal = DataAccessLayer() +Table = TypeVar("Table", bound=Base) class PowerballDraw(Base): @@ -102,3 +92,64 @@ class MegaMillionsDraw(Base): 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()