From 98e872c599a3649cadaf4803ac754f8ddfef3a51 Mon Sep 17 00:00:00 2001 From: Chris Smith Date: Sun, 8 Mar 2026 16:10:44 -0400 Subject: [PATCH] Remove unneeded int casts, update load_dataframe --- lottery_predictor/analyze.py | 24 +++++++++++++++++------- 1 file changed, 17 insertions(+), 7 deletions(-) diff --git a/lottery_predictor/analyze.py b/lottery_predictor/analyze.py index 5302e88..72a8e51 100644 --- a/lottery_predictor/analyze.py +++ b/lottery_predictor/analyze.py @@ -1,3 +1,4 @@ +from datetime import datetime from collections import Counter import pandas as pd @@ -5,9 +6,14 @@ from sqlalchemy import select from data.database import Base, dal, project_variables import util.drawing as drawutil +import util.environment as envutil + +project_variables = envutil.load_environment_variables() -def load_dataframe(table_name: str, use_rule_change_date: bool = True) -> pd.DataFrame: +def load_dataframe( + table_name: str, start_date: datetime | None = None, end_date: datetime | None = None + ) -> pd.DataFrame: """connect to the database and load game data into a pandas dataframe""" dal.connect() @@ -15,10 +21,15 @@ def load_dataframe(table_name: str, use_rule_change_date: bool = True) -> pd.Dat game_dict = drawutil.check_table_name(table_name=table_name) if game_dict: - start_date = game_dict['rule_date'] if use_rule_change_date else None target_table = Base.metadata.tables.get(game_dict['db_table_name']) - if start_date is not None: - sql_statement = select(target_table).where(target_table.columns.draw_date >= start_date) + if start_date is not None and end_date is not None: + sql_statement = (select(target_table) + .where(target_table.columns.draw_date >= start_date) + .where(target_table.columns.draw_date <= end_date)) + elif start_date is not None and end_date is None: + sql_statement = (select(target_table).where(target_table.columns.draw_date >= start_date)) + elif start_date is None and end_date is not None: + sql_statement = (select(target_table).where(target_table.columns.draw_date <= end_date)) else: sql_statement = select(target_table) @@ -36,7 +47,7 @@ def get_most_common_number( columns = ['main_ball1', 'main_ball2', 'main_ball3', 'main_ball4', 'main_ball5'] flat_numbers = data_frame[columns].values.flatten() counts = Counter(flat_numbers) - return [num for num, _ in counts.most_common(top)] + return [int(num) for num, _ in counts.most_common(top)] def get_least_common_number( @@ -48,5 +59,4 @@ def get_least_common_number( columns = ['main_ball1', 'main_ball2', 'main_ball3', 'main_ball4', 'main_ball5'] flat_numbers = data_frame[columns].values.flatten() counts = Counter(flat_numbers) - return list(reversed([num for num, _ in counts.most_common()[-bottom:]])) - + return list(reversed([int(num) for num, _ in counts.most_common()[-bottom:]]))