11"""Common database code used by multiple `covid_hosp` scrapers."""
22
33# standard library
4+ from collections import namedtuple
45from contextlib import contextmanager
56import math
67
1112# first party
1213import delphi .operations .secrets as secrets
1314
15+ Columndef = namedtuple ("Columndef" , "csv_name sql_name dtype" )
1416
1517class Database :
1618
1719 def __init__ (self ,
1820 connection ,
1921 table_name = None ,
2022 columns_and_types = None ,
23+ key_columns = None ,
2124 additional_fields = None ):
2225 """Create a new Database object.
2326
@@ -39,7 +42,11 @@ def __init__(self,
3942 self .table_name = table_name
4043 self .publication_col_name = "issue" if table_name == 'covid_hosp_state_timeseries' else \
4144 'publication_date'
42- self .columns_and_types = columns_and_types
45+ self .columns_and_types = {
46+ c .csv_name : c
47+ for c in (columns_and_types if columns_and_types is not None else [])
48+ }
49+ self .key_columns = key_columns if key_columns is not None else []
4350 self .additional_fields = additional_fields if additional_fields is not None else []
4451
4552 @classmethod
@@ -151,47 +158,49 @@ def insert_dataset(self, publication_date, dataframe):
151158 The dataset.
152159 """
153160 dataframe_columns_and_types = [
154- x for x in self .columns_and_types if x [ 0 ] in dataframe .columns
161+ x for x in self .columns_and_types . values () if x . csv_name in dataframe .columns
155162 ]
163+
164+ def nan_safe_dtype (dtype , value ):
165+ if isinstance (value , float ) and math .isnan (value ):
166+ return None
167+ return dtype (value )
168+
169+ # first convert keys and save the results; we'll need them later
170+ for csv_name in self .key_columns :
171+ dataframe .loc [:, csv_name ] = dataframe [csv_name ].map (self .columns_and_types [csv_name ].dtype )
172+
156173 num_columns = 2 + len (dataframe_columns_and_types ) + len (self .additional_fields )
157174 value_placeholders = ', ' .join (['%s' ] * num_columns )
158- columns = ', ' .join (f'`{ i [ 1 ] } `' for i in dataframe_columns_and_types + self .additional_fields )
175+ columns = ', ' .join (f'`{ i . sql_name } `' for i in dataframe_columns_and_types + self .additional_fields )
159176 sql = f'INSERT INTO `{ self .table_name } ` (`id`, `{ self .publication_col_name } `, { columns } ) ' \
160177 f'VALUES ({ value_placeholders } )'
161178 id_and_publication_date = (0 , publication_date )
162179 with self .new_cursor () as cursor :
163180 for _ , row in dataframe .iterrows ():
164181 values = []
165- for name , _ , dtype in dataframe_columns_and_types :
166- if isinstance (row [name ], float ) and math .isnan (row [name ]):
167- values .append (None )
168- else :
169- values .append (dtype (row [name ]))
182+ for c in dataframe_columns_and_types :
183+ values .append (nan_safe_dtype (c .dtype , row [c .csv_name ]))
170184 cursor .execute (sql ,
171185 id_and_publication_date +
172186 tuple (values ) +
173- tuple (i [ 0 ] for i in self .additional_fields ))
187+ tuple (i . csv_name for i in self .additional_fields ))
174188
175189 # deal with non/seldomly updated columns used like a fk table (if this database needs it)
176190 if hasattr (self , 'AGGREGATE_KEY_COLS' ):
177191 ak_cols = self .AGGREGATE_KEY_COLS
178192
179193 # restrict data to just the key columns and remove duplicate rows
180- ak_data = (dataframe [set (ak_cols + self .KEY_COLS )]
181- .sort_values (self .KEY_COLS )[ak_cols ]
194+ # sort by key columns to ensure that the last ON DUPLICATE KEY overwrite
195+ # uses the most-recent aggregate key information
196+ ak_data = (dataframe [set (ak_cols + self .key_columns )]
197+ .sort_values (self .key_columns )[ak_cols ]
182198 .drop_duplicates ())
183199 # cast types
184- dataframe_typemap = {
185- name : dtype
186- for name , _ , dtype in dataframe_columns_and_types
187- }
188200 for col in ak_cols :
189- def cast_but_sidestep_nans (i ):
190- # not the prettiest, but it works to avoid the NaN values that show up in many columns
191- if isinstance (i , float ) and math .isnan (i ):
192- return None
193- return dataframe_typemap [col ](i )
194- ak_data [col ] = ak_data [col ].apply (cast_but_sidestep_nans )
201+ ak_data [col ] = ak_data [col ].map (
202+ lambda value : nan_safe_dtype (self .columns_and_types [col ].dtype , value )
203+ )
195204 # fix NULLs
196205 ak_data = ak_data .to_numpy (na_value = None ).tolist ()
197206
@@ -204,7 +213,7 @@ def cast_but_sidestep_nans(i):
204213 # use aggregate key table alias
205214 ak_table = self .table_name + '_key'
206215 # assemble full SQL statement
207- ak_insert_sql = f'INSERT INTO `{ ak_table } ` ({ ak_cols_str } ) VALUES ({ values_str } ) as v ON DUPLICATE KEY UPDATE { ak_updates_str } '
216+ ak_insert_sql = f'INSERT INTO `{ ak_table } ` ({ ak_cols_str } ) VALUES ({ values_str } ) AS v ON DUPLICATE KEY UPDATE { ak_updates_str } '
208217
209218 # commit the data
210219 with self .new_cursor () as cur :
0 commit comments