From b497e6d54d793f4143fae055feb3f29a7b081325 Mon Sep 17 00:00:00 2001 From: Frederik Kratzert Date: Wed, 15 Oct 2025 08:26:46 +0000 Subject: [PATCH] get_data now returns time indexed dataframes --- examples/test_australia_fetcher.py | 4 ++-- examples/test_canada_fetcher.py | 4 ++-- examples/test_chile_fetcher.py | 4 ++-- examples/test_france_fetcher.py | 6 +++--- examples/test_japan_fetcher.py | 4 ++-- examples/test_poland_fetcher.py | 4 ++-- examples/test_slovenia_fetcher.py | 4 ++-- examples/test_southafrica_fetcher.py | 4 ++-- examples/test_uk_fetcher.py | 2 +- examples/test_uk_nrfa_fetcher.py | 4 ++-- examples/test_usa_fetcher.py | 4 ++-- rivretrieve/australia.py | 2 +- rivretrieve/base.py | 6 ++++-- rivretrieve/canada.py | 2 +- rivretrieve/chile.py | 4 ++-- rivretrieve/france.py | 2 +- rivretrieve/japan.py | 6 +++--- rivretrieve/poland.py | 2 +- rivretrieve/slovenia.py | 4 ++-- rivretrieve/southafrica.py | 2 +- rivretrieve/uk.py | 4 ++-- rivretrieve/uk_nrfa.py | 4 ++-- rivretrieve/usa.py | 2 +- tests/test_australia.py | 2 +- tests/test_canada.py | 4 ++-- tests/test_chile.py | 2 +- tests/test_france.py | 2 +- tests/test_japan.py | 4 ++-- tests/test_poland.py | 8 ++++---- tests/test_slovenia.py | 8 ++++---- tests/test_uk.py | 4 ++-- tests/test_uk_nrfa.py | 4 ++-- tests/test_usa.py | 4 ++-- 33 files changed, 64 insertions(+), 62 deletions(-) diff --git a/examples/test_australia_fetcher.py b/examples/test_australia_fetcher.py index 122166d..2f5b424 100644 --- a/examples/test_australia_fetcher.py +++ b/examples/test_australia_fetcher.py @@ -23,9 +23,9 @@ if not data.empty: print(f"Data for {gauge_id}:") print(data.head()) - print(f"Time series from {data[constants.TIME_INDEX].min()} to {data[constants.TIME_INDEX].max()}") + print(f"Time series from {data.index.min()} to {data.index.max()}") plt.plot( - data[constants.TIME_INDEX], + data.index, data[constants.DISCHARGE], label=gauge_id, marker="o", diff --git a/examples/test_canada_fetcher.py b/examples/test_canada_fetcher.py index 29abfc4..7e6b7b6 100644 --- a/examples/test_canada_fetcher.py +++ b/examples/test_canada_fetcher.py @@ -20,9 +20,9 @@ if not data.empty: print(f"Data for {gauge_id}:") print(data.head()) - print(f"Time series from {data[constants.TIME_INDEX].min()} to {data[constants.TIME_INDEX].max()}") + print(f"Time series from {data.index.min()} to {data.index.max()}") plt.plot( - data[constants.TIME_INDEX], + data.index, data[constants.DISCHARGE], label=gauge_id, marker=".", diff --git a/examples/test_chile_fetcher.py b/examples/test_chile_fetcher.py index 4961d2a..9552375 100644 --- a/examples/test_chile_fetcher.py +++ b/examples/test_chile_fetcher.py @@ -20,9 +20,9 @@ if not data.empty: print(f"Data for {gauge_id}:") print(data.head()) - print(f"Time series from {data[constants.TIME_INDEX].min()} to {data[constants.TIME_INDEX].max()}") + print(f"Time series from {data.index.min()} to {data.index.max()}") plt.plot( - data[constants.TIME_INDEX], + data.index, data[constants.DISCHARGE], label=gauge_id, marker=".", diff --git a/examples/test_france_fetcher.py b/examples/test_france_fetcher.py index e5fbb5d..da1983b 100644 --- a/examples/test_france_fetcher.py +++ b/examples/test_france_fetcher.py @@ -20,15 +20,15 @@ if not data.empty: print(f"Data for {gauge_id}:") print(data.head()) - print(f"Time series from {data[constants.TIME_INDEX].min()} to {data[constants.TIME_INDEX].max()}") + print(f"Time series from {data.index.min()} to {data.index.max()}") plt.plot( - data[constants.TIME_INDEX], + data.index, data[constants.DISCHARGE], label=gauge_id, marker=".", linestyle="-", ) - plt.xlim(data[constants.TIME_INDEX].min(), data[constants.TIME_INDEX].max()) + plt.xlim(data.index.min(), data.index.max()) else: print(f"No data found for {gauge_id}") diff --git a/examples/test_japan_fetcher.py b/examples/test_japan_fetcher.py index a9c573b..ca565cc 100644 --- a/examples/test_japan_fetcher.py +++ b/examples/test_japan_fetcher.py @@ -22,9 +22,9 @@ if not data.empty: print(f"Data for {gauge_id}:") print(data.head()) - print(f"Time series from {data[constants.TIME_INDEX].min()} to {data[constants.TIME_INDEX].max()}") + print(f"Time series from {data.index.min()} to {data.index.max()}") plt.plot( - data[constants.TIME_INDEX], + data.index, data[constants.DISCHARGE], label=gauge_id, marker="o", diff --git a/examples/test_poland_fetcher.py b/examples/test_poland_fetcher.py index 6df1e6b..782c726 100644 --- a/examples/test_poland_fetcher.py +++ b/examples/test_poland_fetcher.py @@ -29,9 +29,9 @@ if not data.empty: print(f"Data for {gauge_id}:") print(data.head()) - print(f"Time series from {data[constants.TIME_INDEX].min()} to {data[constants.TIME_INDEX].max()}") + print(f"Time series from {data.index.min()} to {data.index.max()}") plt.plot( - data[constants.TIME_INDEX], + data.index, data[variable], label=gauge_id, marker=".", diff --git a/examples/test_slovenia_fetcher.py b/examples/test_slovenia_fetcher.py index 65ff03c..c3a822c 100644 --- a/examples/test_slovenia_fetcher.py +++ b/examples/test_slovenia_fetcher.py @@ -20,9 +20,9 @@ if not data.empty: print(f"Data for {gauge_id}:") print(data.head()) - print(f"Time series from {data[constants.TIME_INDEX].min()} to {data[constants.TIME_INDEX].max()}") + print(f"Time series from {data.index.min()} to {data.index.max()}") plt.plot( - data[constants.TIME_INDEX], + data.index, data[constants.DISCHARGE], label=gauge_id, marker="o", diff --git a/examples/test_southafrica_fetcher.py b/examples/test_southafrica_fetcher.py index 23cb818..3a98f40 100644 --- a/examples/test_southafrica_fetcher.py +++ b/examples/test_southafrica_fetcher.py @@ -20,9 +20,9 @@ if not data.empty: print(f"Data for {gauge_id}:") print(data.head()) - print(f"Time series from {data[constants.TIME_INDEX].min()} to {data[constants.TIME_INDEX].max()}") + print(f"Time series from {data.index.min()} to {data.index.max()}") plt.plot( - data[constants.TIME_INDEX], + data.index, data[constants.DISCHARGE], label=gauge_id, marker=".", diff --git a/examples/test_uk_fetcher.py b/examples/test_uk_fetcher.py index 2fef0a1..19031eb 100644 --- a/examples/test_uk_fetcher.py +++ b/examples/test_uk_fetcher.py @@ -19,7 +19,7 @@ print(f"Data for {gauge_id}:") print(data.head()) plt.plot( - data[constants.TIME_INDEX], + data.index, data[constants.DISCHARGE], label=gauge_id.split("/")[-1], ) diff --git a/examples/test_uk_nrfa_fetcher.py b/examples/test_uk_nrfa_fetcher.py index 14d76c9..54c6894 100644 --- a/examples/test_uk_nrfa_fetcher.py +++ b/examples/test_uk_nrfa_fetcher.py @@ -31,9 +31,9 @@ if not data.empty: print(f"Data for {gauge_id}:") print(data.head()) - print(f"Time series from {data[constants.TIME_INDEX].min()} to {data[constants.TIME_INDEX].max()}") + print(f"Time series from {data.index.min()} to {data.index.max()}") plt.plot( - data[constants.TIME_INDEX], + data.index, data[constants.DISCHARGE], label=gauge_id, marker=".", diff --git a/examples/test_usa_fetcher.py b/examples/test_usa_fetcher.py index 4d12573..da6ba17 100644 --- a/examples/test_usa_fetcher.py +++ b/examples/test_usa_fetcher.py @@ -23,9 +23,9 @@ if not data.empty: print(f"Data for {gauge_id}:") print(data.head()) - print(f"Time series from {data[constants.TIME_INDEX].min()} to {data[constants.TIME_INDEX].max()}") + print(f"Time series from {data.index.min()} to {data.index.max()}") plt.plot( - data[constants.TIME_INDEX], + data.index, data[constants.DISCHARGE], label=gauge_id, marker="o", diff --git a/rivretrieve/australia.py b/rivretrieve/australia.py index 01525ab..4a1b445 100644 --- a/rivretrieve/australia.py +++ b/rivretrieve/australia.py @@ -148,7 +148,7 @@ def _parse_data(self, gauge_id: str, raw_data: Optional[str], variable: str) -> df["Value"] = pd.to_numeric(df["Value"], errors="coerce") df = df.rename(columns={"Value": variable}) df[constants.TIME_INDEX] = pd.to_datetime(df[constants.TIME_INDEX]) - return df[[constants.TIME_INDEX, variable]].dropna() + return df[[constants.TIME_INDEX, variable]].dropna().set_index(constants.TIME_INDEX) except Exception as e: logger.error(f"Error parsing CSV data for site {gauge_id}: {e}") return pd.DataFrame(columns=[constants.TIME_INDEX, variable]) diff --git a/rivretrieve/base.py b/rivretrieve/base.py index 7ea8010..fd6746c 100644 --- a/rivretrieve/base.py +++ b/rivretrieve/base.py @@ -25,12 +25,14 @@ def get_data( Args: gauge_id: The site-specific identifier for the gauge. - variable: The variable to fetch (e.g., 'stage' or 'discharge'). + variable: The variable to fetch, should be one of the values from constants.py + (e.g., constants.DISCHARGE, constants.STAGE). start_date: Optional start date in 'YYYY-MM-DD' format. end_date: Optional end date in 'YYYY-MM-DD' format. Returns: - A pandas DataFrame with 'Date' and the variable column ('H' or 'Q'). + A pandas DataFrame indexed by time (constants.TIME_INDEX) with a column + for the requested variable (e.g., constants.DISCHARGE). """ pass diff --git a/rivretrieve/canada.py b/rivretrieve/canada.py index 9d788a5..589c414 100644 --- a/rivretrieve/canada.py +++ b/rivretrieve/canada.py @@ -186,7 +186,7 @@ def get_data( df_long[[constants.TIME_INDEX, variable]] .dropna() .sort_values(by=constants.TIME_INDEX) - .reset_index(drop=True) + .set_index(constants.TIME_INDEX) ) except Exception as e: diff --git a/rivretrieve/chile.py b/rivretrieve/chile.py index 930df5e..0ad25de 100644 --- a/rivretrieve/chile.py +++ b/rivretrieve/chile.py @@ -104,7 +104,7 @@ def _parse_data( df[[constants.TIME_INDEX, variable]] .dropna() .sort_values(by=constants.TIME_INDEX) - .reset_index(drop=True) + .set_index(constants.TIME_INDEX) ) except Exception as e: logger.error(f"Error parsing data for site {gauge_id}: {e}") @@ -135,7 +135,7 @@ def get_data( # Filter by date range start_date_dt = pd.to_datetime(start_date) end_date_dt = pd.to_datetime(end_date) - df = df[(df[constants.TIME_INDEX] >= start_date_dt) & (df[constants.TIME_INDEX] <= end_date_dt)] + df = df[(df.index >= start_date_dt) & (df.index <= end_date_dt)] return df except Exception as e: logger.error(f"Failed to get data for site {gauge_id}, variable {variable}: {e}") diff --git a/rivretrieve/france.py b/rivretrieve/france.py index 96972b9..64791fa 100644 --- a/rivretrieve/france.py +++ b/rivretrieve/france.py @@ -107,7 +107,7 @@ def _parse_data( df[[constants.TIME_INDEX, variable]] .dropna() .sort_values(by=constants.TIME_INDEX) - .reset_index(drop=True) + .set_index(constants.TIME_INDEX) ) except Exception as e: logger.error(f"Error parsing JSON data for site {gauge_id}: {e}") diff --git a/rivretrieve/japan.py b/rivretrieve/japan.py index 8817dcb..eb818c2 100644 --- a/rivretrieve/japan.py +++ b/rivretrieve/japan.py @@ -141,9 +141,9 @@ def _parse_data( final_df = pd.concat(all_dfs, ignore_index=True) final_df = final_df.rename(columns={"Value": variable}) - final_df = final_df.sort_values(by=constants.TIME_INDEX).reset_index(drop=True) + final_df = final_df.sort_values(by=constants.TIME_INDEX) - return final_df + return final_df.set_index(constants.TIME_INDEX) def get_data( self, @@ -164,7 +164,7 @@ def get_data( start_date_dt = pd.to_datetime(start_date) end_date_dt = pd.to_datetime(end_date) - df = df[(df[constants.TIME_INDEX] >= start_date_dt) & (df[constants.TIME_INDEX] <= end_date_dt)] + df = df[(df.index >= start_date_dt) & (df.index <= end_date_dt)] return df except Exception as e: diff --git a/rivretrieve/poland.py b/rivretrieve/poland.py index 1f46116..f8619e8 100644 --- a/rivretrieve/poland.py +++ b/rivretrieve/poland.py @@ -209,7 +209,7 @@ def get_data( data_array = ds[variable].sel(gauge_id=gauge_id, time=slice(start_date, end_date)) df = data_array.to_pandas().dropna().reset_index().rename(columns={variable: variable}) - return df[[constants.TIME_INDEX, variable]] + return df.set_index(constants.TIME_INDEX)[[variable]] except KeyError: logger.info(f"No data found for gauge {gauge_id} in the selected date range.") diff --git a/rivretrieve/slovenia.py b/rivretrieve/slovenia.py index 72f8159..4fb0f44 100644 --- a/rivretrieve/slovenia.py +++ b/rivretrieve/slovenia.py @@ -85,7 +85,7 @@ def _parse_data(self, gauge_id: str, raw_data: Optional[str], variable: str) -> df[[constants.TIME_INDEX, variable]] .dropna() .sort_values(by=constants.TIME_INDEX) - .reset_index(drop=True) + .set_index(constants.TIME_INDEX) ) except Exception as e: @@ -113,7 +113,7 @@ def get_data( # Filter by date range start_date_dt = pd.to_datetime(start_date) end_date_dt = pd.to_datetime(end_date) - df = df[(df[constants.TIME_INDEX] >= start_date_dt) & (df[constants.TIME_INDEX] <= end_date_dt)] + df = df[(df.index >= start_date_dt) & (df.index <= end_date_dt)] return df except Exception as e: logger.error(f"Failed to get data for site {gauge_id}, variable {variable}: {e}") diff --git a/rivretrieve/southafrica.py b/rivretrieve/southafrica.py index 6d403b6..450e8ac 100644 --- a/rivretrieve/southafrica.py +++ b/rivretrieve/southafrica.py @@ -155,7 +155,7 @@ def _parse_data( daily_df = full_df[[constants.TIME_INDEX, "D_AVG_FR"]].rename(columns={"D_AVG_FR": "Value"}) daily_df = daily_df.rename(columns={"Value": variable}) - return daily_df.dropna().sort_values(by=constants.TIME_INDEX).reset_index(drop=True) + return daily_df.dropna().sort_values(by=constants.TIME_INDEX).set_index(constants.TIME_INDEX) except Exception as e: logger.error(f"Error parsing data for site {gauge_id}: {e}") diff --git a/rivretrieve/uk.py b/rivretrieve/uk.py index 3754275..237315d 100644 --- a/rivretrieve/uk.py +++ b/rivretrieve/uk.py @@ -126,7 +126,7 @@ def _parse_data(self, gauge_id: str, raw_data: List[Dict[str, Any]], variable: s complete_ts = pd.DataFrame(date_range, columns=[constants.TIME_INDEX]) df_daily = pd.merge(complete_ts, df_daily, on=constants.TIME_INDEX, how="left") - return df_daily + return df_daily.set_index(constants.TIME_INDEX) def get_data( self, @@ -148,7 +148,7 @@ def get_data( # Filter by exact start and end date after processing start_date_dt = pd.to_datetime(start_date) end_date_dt = pd.to_datetime(end_date) - df = df[(df[constants.TIME_INDEX] >= start_date_dt) & (df[constants.TIME_INDEX] <= end_date_dt)] + df = df[(df.index >= start_date_dt) & (df.index <= end_date_dt)] return df except Exception as e: diff --git a/rivretrieve/uk_nrfa.py b/rivretrieve/uk_nrfa.py index 07186c5..18747a4 100644 --- a/rivretrieve/uk_nrfa.py +++ b/rivretrieve/uk_nrfa.py @@ -102,7 +102,7 @@ def _parse_data(self, gauge_id: str, raw_data: Optional[Dict[str, Any]], variabl df[constants.TIME_INDEX] = pd.to_datetime(df["time"], format="ISO8601").dt.date df[constants.TIME_INDEX] = pd.to_datetime(df[constants.TIME_INDEX]) df[variable] = pd.to_numeric(df[variable], errors="coerce") - return df[[constants.TIME_INDEX, variable]].dropna().reset_index(drop=True) + return df[[constants.TIME_INDEX, variable]].dropna().set_index(constants.TIME_INDEX) except Exception as e: logger.error(f"Error parsing NRFA data for {gauge_id}: {e}") return pd.DataFrame(columns=[constants.TIME_INDEX, variable]) @@ -128,7 +128,7 @@ def get_data( # Filter by date range start_date_dt = pd.to_datetime(start_date) end_date_dt = pd.to_datetime(end_date) - df = df[(df[constants.TIME_INDEX] >= start_date_dt) & (df[constants.TIME_INDEX] <= end_date_dt)] + df = df[(df.index >= start_date_dt) & (df.index <= end_date_dt)] return df except Exception as e: logger.error(f"Failed to get data for site {gauge_id}, variable {variable}: {e}") diff --git a/rivretrieve/usa.py b/rivretrieve/usa.py index c689fc1..0ded443 100644 --- a/rivretrieve/usa.py +++ b/rivretrieve/usa.py @@ -76,7 +76,7 @@ def _parse_data(self, gauge_id: str, raw_data: pd.DataFrame, variable: str) -> p mult = 0.0283168466 df[variable] = pd.to_numeric(df[value_col], errors="coerce") * mult - return df[[constants.TIME_INDEX, variable]].dropna() + return df[[constants.TIME_INDEX, variable]].dropna().set_index(constants.TIME_INDEX) def get_data( self, diff --git a/tests/test_australia.py b/tests/test_australia.py index a81c9d6..9f63588 100644 --- a/tests/test_australia.py +++ b/tests/test_australia.py @@ -47,7 +47,7 @@ def bom_request_side_effect(params): constants.TIME_INDEX: pd.to_datetime(["2010-01-01", "2010-01-02", "2010-01-03"]), constants.DISCHARGE: [0.000, 3.710, 3.211], } - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) assert_frame_equal(result_df, expected_df) self.assertEqual(mock_make_bom_request.call_count, 2) diff --git a/tests/test_canada.py b/tests/test_canada.py index 6771741..85c2079 100644 --- a/tests/test_canada.py +++ b/tests/test_canada.py @@ -36,7 +36,7 @@ def test_get_data_discharge(self, mock_hydat_path, mock_download, mock_requests) ), constants.DISCHARGE: [1.1, 1.2, 1.3, 1.4, 1.5], } - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) assert_frame_equal(result_df, expected_df) @@ -62,7 +62,7 @@ def test_get_data_stage(self, mock_hydat_path, mock_download, mock_requests): ), constants.STAGE: [10.1, 10.2, 10.3, 10.4, 10.5], } - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) assert_frame_equal(result_df, expected_df) diff --git a/tests/test_chile.py b/tests/test_chile.py index 0f35af0..cd86075 100644 --- a/tests/test_chile.py +++ b/tests/test_chile.py @@ -50,7 +50,7 @@ def get_side_effect(*args, **kwargs): constants.TIME_INDEX: pd.to_datetime(["2022-01-01", "2022-01-02", "2022-01-03"]), constants.DISCHARGE: [15.5, 16.0, 15.8], } - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) assert_frame_equal(result_df, expected_df) self.assertEqual(mock_get.call_count, 2) diff --git a/tests/test_france.py b/tests/test_france.py index 9d056d7..6d5701e 100644 --- a/tests/test_france.py +++ b/tests/test_france.py @@ -38,7 +38,7 @@ def test_get_data_discharge(self, mock_get): constants.TIME_INDEX: pd.to_datetime(["2023-01-01", "2023-01-02", "2023-01-03"]), constants.DISCHARGE: [15.0005, 16.0000, 15.5002], # Divided by 1000 } - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) assert_frame_equal(result_df, expected_df) mock_get.assert_called_once() diff --git a/tests/test_japan.py b/tests/test_japan.py index 7501af3..423c65b 100644 --- a/tests/test_japan.py +++ b/tests/test_japan.py @@ -40,9 +40,9 @@ def test_get_data_discharge(self, mock_get): constants.TIME_INDEX: expected_dates, constants.DISCHARGE: expected_values, } - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) - assert_frame_equal(result_df.reset_index(drop=True), expected_df) + assert_frame_equal(result_df, expected_df) mock_get.assert_called_once() # Check that the params are correct mock_args, mock_kwargs = mock_get.call_args diff --git a/tests/test_poland.py b/tests/test_poland.py index 489f564..22e1681 100644 --- a/tests/test_poland.py +++ b/tests/test_poland.py @@ -46,9 +46,9 @@ def test_get_data_discharge(self, mock_create_cache): constants.TIME_INDEX: expected_dates, constants.DISCHARGE: expected_values, } - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) - assert_frame_equal(result_df.reset_index(drop=True), expected_df) + assert_frame_equal(result_df, expected_df) mock_create_cache.assert_not_called() # Cache should not be recreated @patch("rivretrieve.poland.PolandFetcher._create_cache") # Mock cache creation @@ -67,9 +67,9 @@ def test_get_data_stage(self, mock_create_cache): constants.TIME_INDEX: expected_dates, constants.STAGE: expected_values, } - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) - assert_frame_equal(result_df.reset_index(drop=True), expected_df) + assert_frame_equal(result_df, expected_df) mock_create_cache.assert_not_called() @patch("rivretrieve.utils.requests_retry_session") diff --git a/tests/test_slovenia.py b/tests/test_slovenia.py index 1c66615..c254f29 100644 --- a/tests/test_slovenia.py +++ b/tests/test_slovenia.py @@ -45,9 +45,9 @@ def test_get_data_discharge(self, mock_requests_session): constants.TIME_INDEX: expected_dates, constants.DISCHARGE: expected_values, } - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) - assert_frame_equal(result_df.reset_index(drop=True), expected_df) + assert_frame_equal(result_df, expected_df) mock_session.get.assert_called_once() mock_args, mock_kwargs = mock_session.get.call_args self.assertIn("p_postaja=1020", mock_args[0]) @@ -81,9 +81,9 @@ def test_get_data_stage(self, mock_requests_session): constants.TIME_INDEX: expected_dates, constants.STAGE: expected_values, } - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) - assert_frame_equal(result_df.reset_index(drop=True), expected_df) + assert_frame_equal(result_df, expected_df) if __name__ == "__main__": diff --git a/tests/test_uk.py b/tests/test_uk.py index ec40aeb..ab980cd 100644 --- a/tests/test_uk.py +++ b/tests/test_uk.py @@ -61,9 +61,9 @@ def mock_get_side_effect(url, *args, **kwargs): constants.TIME_INDEX: expected_dates, constants.DISCHARGE: expected_values, } - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) - assert_frame_equal(result_df.reset_index(drop=True), expected_df, check_dtype=False) + assert_frame_equal(result_df, expected_df, check_dtype=False) self.assertEqual(mock_session.get.call_count, 2) diff --git a/tests/test_uk_nrfa.py b/tests/test_uk_nrfa.py index 6e5aa08..56bc1d9 100644 --- a/tests/test_uk_nrfa.py +++ b/tests/test_uk_nrfa.py @@ -99,9 +99,9 @@ def test_get_data(self, variable, expected_data_type, sample_file, expected_valu expected_dates = pd.to_datetime(["2022-01-01", "2022-01-02", "2022-01-03", "2022-01-04", "2022-01-05"]) expected_data = {constants.TIME_INDEX: expected_dates, variable: expected_values} - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) - assert_frame_equal(result_df.reset_index(drop=True), expected_df, check_dtype=False) + assert_frame_equal(result_df, expected_df, check_dtype=False) mock_session.get.assert_called_once() _, mock_kwargs = mock_session.get.call_args params = mock_kwargs["params"] diff --git a/tests/test_usa.py b/tests/test_usa.py index 09792d2..e7a16a2 100644 --- a/tests/test_usa.py +++ b/tests/test_usa.py @@ -51,9 +51,9 @@ def test_get_data_discharge(self, mock_get_dv): constants.TIME_INDEX: expected_dates, constants.DISCHARGE: expected_values, } - expected_df = pd.DataFrame(expected_data) + expected_df = pd.DataFrame(expected_data).set_index(constants.TIME_INDEX) - assert_frame_equal(result_df.reset_index(drop=True), expected_df, check_dtype=False) + assert_frame_equal(result_df, expected_df, check_dtype=False) mock_get_dv.assert_called_once() mock_args, mock_kwargs = mock_get_dv.call_args self.assertEqual(mock_kwargs["sites"], gauge_id)