Add new DAL and update code to create missing database
This commit is contained in:
1 parent
bbdb90c7e2
commit
3fe2e1bcd5
1 file changed
+68
-17
+68
-17
@@ -3,10 +3,13 @@
|
|||||||
|
|
||||||
The base definitions for database access and definition
|
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 import create_engine, Date, engine, func, Integer, MetaData, Row, text
|
||||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, sessionmaker
|
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
|
from util.environment import load_environment_variables
|
||||||
|
|
||||||
@@ -32,20 +35,7 @@ class Base(DeclarativeBase):
|
|||||||
metadata = MetaData(naming_convention=DATABASE_NAMING_CONVENTION)
|
metadata = MetaData(naming_convention=DATABASE_NAMING_CONVENTION)
|
||||||
|
|
||||||
|
|
||||||
class DataAccessLayer:
|
Table = TypeVar("Table", bound=Base)
|
||||||
|
|
||||||
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()
|
|
||||||
|
|
||||||
|
|
||||||
class PowerballDraw(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_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.main_ball4 == other.main_ball4 and self.main_ball5 == other.main_ball5 and \
|
||||||
self.mega_ball == other.mega_ball and self.megaplier == other.megaplier
|
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()
|
||||||
Reference in new issue
Block a user