datajoi/Dataset-Test-Workflow
3
1import os2import json3import duckdb4import gradio as gr5import pandas as pd6import pandera as pa7from pandera import Column8import ydata_profiling as pp9from langchain_huggingface import HuggingFaceEndpoint, ChatHuggingFace10from langsmith import traceable11from langchain import hub12import warnings13import dlt14warnings.filterwarnings("ignore", category=DeprecationWarning)15 16# Height of the Tabs Text Area17TAB_LINES = 818 19 20#----------CONNECT TO DATABASE----------21md_token = os.getenv('MD_TOKEN')22conn = duckdb.connect(f"md:my_db?motherduck_token={md_token}", read_only=True)23#---------------------------------------24 25 26#-------LOAD HUGGINGFACE-------27models = ["Qwen/Qwen2.5-72B-Instruct","meta-llama/Meta-Llama-3-70B-Instruct",28 "meta-llama/Llama-3.1-70B-Instruct"]29 30model_loaded = False 31for model in models:32 try:33 endpoint = HuggingFaceEndpoint(repo_id=model, max_new_tokens=8192)34 info = endpoint.client.get_endpoint_info()35 model_loaded = True36 break37 except Exception as e:38 print(f"Error for model {model}: {e}")39 continue40 41llm = ChatHuggingFace(llm=endpoint).bind(max_tokens=8192)42#---------------------------------------43 44#-----LOAD PROMPT FROM LANCHAIN HUB-----45prompt_autogenerate = hub.pull("autogenerate-rules-testworkflow")46prompt_user_input = hub.pull("usergenerate-rules-testworkflow")47 48#--------------ALL UTILS----------------49# Get Databases50def get_schemas():51 schemas = conn.execute("""52 SELECT DISTINCT schema_name53 FROM information_schema.schemata54 WHERE schema_name NOT IN ('information_schema', 'pg_catalog')55 """).fetchall()56 return [item[0] for item in schemas]57 58# Get Tables59def get_tables_names(schema_name):60 tables = conn.execute(f"SELECT table_name FROM information_schema.tables WHERE table_schema = '{schema_name}'").fetchall()61 return [table[0] for table in tables]62 63# Update Tables64def update_table_names(schema_name):65 tables = get_tables_names(schema_name)66 return gr.update(choices=tables)67# def get_data_df(schema):68# print('Getting Dataframe from the Database')69# return conn.sql(f"SELECT * FROM {schema} LIMIT 1000")70 71@dlt.resource72def fetch_data(schema):73 result = conn.sql(f"SELECT * FROM {schema} LIMIT 1000")74 75 while True:76 chunk_df = result.fetch_df_chunk(2)77 78 if chunk_df is None or len(chunk_df) == 0:79 break80 else:81 yield chunk_df82 83def create_pipeline(schema):84 dataset_name = schema.split('.')[1]85 print("Dataset Name: ", dataset_name)86 87 table_name = schema.split('.')[2]88 print("Table Name: ", table_name)89 90 pipeline =dlt.pipeline(91 pipeline_name='duckdb_pipeline',92 destination='duckdb',93 dataset_name= dataset_name,94 )95 96 load_info = pipeline.run(fetch_data(schema), table_name = table_name,97 write_disposition = "replace")98 99 print(load_info)100 return dataset_name + "." + table_name 101 102def load_pipeline(table_name):103 _conn = duckdb.connect("duckdb_pipeline.duckdb")104 return _conn, _conn.sql(f"SELECT * FROM {table_name} LIMIT 1000").df()105 106def df_summary(df):107 summary = []108 109 for column in df.columns:110 if pd.api.types.is_numeric_dtype(df[column]):111 summary.append({112 "column": column,113 "max": df[column].max(),114 "min": df[column].min(),115 "count": df[column].count(),116 "nunique": df[column].nunique(),117 "dtype": str(df[column].dtype),118 "top": None119 })120 121 elif pd.api.types.is_categorical_dtype(df[column]) or pd.api.types.is_object_dtype(df[column]):122 top_value = df[column].mode().iloc[0] if not df[column].mode().empty else None123 124 summary.append({125 "column": column,126 "max": None, 127 "min": None, 128 "count": df[column].count(),129 "nunique": df[column].nunique(),130 "dtype": str(df[column].dtype),131 "top": top_value132 })133 summary_df = pd.DataFrame(summary)134 return summary_df.reset_index(drop=True)135 136def format_prompt(df):137 summary = df_summary(df)138 return prompt_autogenerate.format_prompt(data=df.head().to_json(orient='records'),139 summary=summary.to_json(orient='records'))140def format_user_prompt(df):141 return prompt_user_input.format_prompt(data=df.head().to_json(orient='records'))142 143def process_inputs(inputs) :144 return {'input_query': inputs['messages'].to_messages()[1]}145 146@traceable(process_inputs=process_inputs)147def run_llm(messages):148 try:149 response = llm.invoke(messages)150 print(response.content.replace("```", "'''").replace("json", ""))151 tests = json.loads(response.content.replace("```", "").replace("json", ""))152 except Exception as e:153 return e154 return tests155 156 157# Get Schema158def get_table_schema(table):159 result = conn.sql(f"SELECT sql, database_name, schema_name FROM duckdb_tables() where table_name ='{table}';").df()160 ddl_create = result.iloc[0,0]161 parent_database = result.iloc[0,1]162 schema_name = result.iloc[0,2]163 full_path = f"{parent_database}.{schema_name}.{table}"164 if schema_name != "main":165 old_path = f"{schema_name}.{table}"166 else:167 old_path = table168 ddl_create = ddl_create.replace(old_path, full_path)169 return full_path170 171def describe(df):172 173 numerical_info = pd.DataFrame()174 categorical_info = pd.DataFrame()175 if len(df.select_dtypes(include=['number']).columns) >= 1:176 numerical_info = df.select_dtypes(include=['number']).describe().T.reset_index()177 numerical_info.rename(columns={'index': 'column'}, inplace=True)178 if len(df.select_dtypes(include=['object']).columns) >= 1:179 categorical_info = df.select_dtypes(include=['object']).describe().T.reset_index()180 categorical_info.rename(columns={'index': 'column'}, inplace=True)181 182 return numerical_info, categorical_info183 184def validate_pandera(tests, df):185 validation_results = []186 187 for test in tests:188 column_name = test['column_name']189 try:190 rule = eval(test['pandera_rule']) 191 validated_column = rule(df[[column_name]]) 192 validation_results.append({193 "Columns": column_name,194 "Result": "✅ Pass"195 })196 except Exception as e:197 validation_results.append({198 "Columns": column_name,199 "Result": f"❌ Fail - {str(e)}"200 })201 return pd.DataFrame(validation_results)202 203def statistics(df):204 profile = pp.ProfileReport(df)205 report_dict = profile.get_description()206 description, alerts = report_dict.table, report_dict.alerts207 # Statistics208 mapping = {209 'n': 'Number of observations',210 'n_var': 'Number of variables',211 'n_cells_missing': 'Number of cells missing',212 'n_vars_with_missing': 'Number of columns with missing data',213 'n_vars_all_missing': 'Columns with all missing data',214 'p_cells_missing': 'Missing cells (%)',215 'n_duplicates': 'Duplicated rows',216 'p_duplicates': 'Duplicated rows (%)',217 }218 219 updated_data = {mapping.get(k, k): v for k, v in description.items() if k != 'types'}220 # Add flattened types information221 if 'Text' in description.get('types', {}):222 updated_data['Number of text columns'] = description['types']['Text']223 if 'Categorical' in description.get('types', {}):224 updated_data['Number of categorical columns'] = description['types']['Categorical']225 if 'Numeric' in description.get('types', {}):226 updated_data['Number of numeric columns'] = description['types']['Numeric']227 if 'DateTime' in description.get('types', {}):228 updated_data['Number of datetime columns'] = description['types']['DateTime']229 230 df_statistics = pd.DataFrame(list(updated_data.items()), columns=['Statistic Description', 'Value'])231 df_statistics['Value'] = df_statistics['Value'].astype(int)232 233 # Alerts234 alerts_list = [(str(alert).replace('[', '').replace(']', ''), alert.alert_type_name) for alert in alerts]235 df_alerts = pd.DataFrame(alerts_list, columns=['Data Quality Issue', 'Category'])236 237 return df_statistics, df_alerts238#---------------------------------------239 240 241 242# Main Function243def main(table):244 schema = get_table_schema(table)245 246 # Create dlt pipeline247 table_name = create_pipeline(schema)248 249 # Load dlt pipeline250 connection, df = load_pipeline(table_name)251 252 # df = get_data_df(schema)253 df_statistics, df_alerts = statistics(df)254 describe_num, describe_cat = describe(df)255 256 messages = format_prompt(df=df)257 tests = run_llm(messages)258 259 if isinstance(tests, Exception):260 tests = pd.DataFrame([{"error": f"❌ Unable to generate tests. {tests}"}])261 return df.head(10), df_statistics, df_alerts, describe_cat, describe_num, tests, pd.DataFrame([])262 263 tests_df = pd.DataFrame(tests)264 tests_df.rename(columns={tests_df.columns[0]: 'Column', tests_df.columns[1]: 'Rule Name', tests_df.columns[2]: 'Rules' }, inplace=True)265 pandera_results = validate_pandera(tests, df)266 267 connection.close()268 return df.head(10), df_statistics, df_alerts, describe_cat, describe_num, tests_df, pandera_results269 270def user_results(table, text_query):271 272 schema = get_table_schema(table)273 274 # Create dlt pipeline275 table_name = create_pipeline(schema)276 277 # Load dlt pipeline278 connection, df = load_pipeline(table_name)279 280 messages = format_user_prompt(df=df, user_description=text_query)281 282 print(f'Generated Tests from user input: {tests}')283 284 if isinstance(tests, Exception):285 tests = pd.DataFrame([{"error": f"❌ Unable to generate tests. {tests}"}])286 return tests, pd.DataFrame([])287 288 tests_df = pd.DataFrame(tests)289 tests_df.rename(columns={tests_df.columns[0]: 'Column', tests_df.columns[1]: 'Rule Name', tests_df.columns[2]: 'Rules' }, inplace=True)290 pandera_results = validate_pandera(tests, df)291 292 connection.close()293 294 return tests_df, pandera_results295 296# Custom CSS styling297custom_css = """298 print('Validated Tests with Pandera')299.gradio-container {300 background-color: #f0f4f8;301 302}303.logo {304 max-width: 200px;305 margin: 20px auto;306 display: block;307}308.gr-button {309 background-color: #4a90e2 !important;310}311.gr-button:hover {312 background-color: #3a7bc8 !important;313}314"""315 316with gr.Blocks(theme=gr.themes.Soft(primary_hue="purple", secondary_hue="indigo"), css=custom_css) as demo:317 gr.Image("logo.png", label=None, show_label=False, container=False, height=100)318 319 gr.Markdown("""320 <div style='text-align: center;'>321 <strong style='font-size: 36px;'>Dataset Test Workflow</strong>322 <br>323 <span style='font-size: 20px;'>Implement and Automate Data Validation Processes.</span>324 </div>325 """)326 327 with gr.Row():328 with gr.Column(scale=1):329 schema_dropdown = gr.Dropdown(choices=get_schemas(), label="Select Schema", interactive=True)330 tables_dropdown = gr.Dropdown(choices=[], label="Available Tables", value=None)331 with gr.Row():332 generate_result = gr.Button("Validate Data", variant="primary")333 334 with gr.Column(scale=2):335 with gr.Tabs():336 337 with gr.Tab("Description"):338 with gr.Row():339 with gr.Column():340 data_description = gr.DataFrame(label="Data Description", value=[], interactive=False)341 with gr.Row():342 with gr.Column():343 describe_cat = gr.DataFrame(label="Categorical Information", value=[], interactive=False)344 with gr.Column(): 345 describe_num = gr.DataFrame(label="Numerical Information", value=[], interactive=False)346 347 with gr.Tab("Alerts"):348 data_alerts = gr.DataFrame(label="Alerts", value=[], interactive=False)349 350 with gr.Tab("Rules & Validations"):351 tests_output = gr.DataFrame(label="Validation Rules", value=[], interactive=False)352 test_result_output = gr.DataFrame(label="Validation Result", value=[], interactive=False)353 354 with gr.Tab("Data"):355 result_output = gr.DataFrame(label="Dataframe (10 Rows)", value=[], interactive=False)356 357 with gr.Tab('Text to Validation'):358 with gr.Row():359 query_input = gr.Textbox(lines=5, label="Text Query", placeholder="Enter Text Query to Generate Validation e.g. Validate that the incident_zip column contains valid 5-digit ZIP codes.")360 with gr.Row():361 with gr.Column(): 362 pass363 with gr.Column(scale=1, min_width=50): 364 user_generate_result = gr.Button("Validate Data", variant="primary" ) 365 366 with gr.Row():367 with gr.Column():368 query_tests = gr.DataFrame(label="Validation Rules", value=[], interactive=False)369 with gr.Column():370 query_result = gr.DataFrame(label="Validation Result", value=[], interactive=False)371 372 schema_dropdown.change(update_table_names, inputs=schema_dropdown, outputs=tables_dropdown)373 generate_result.click(main, inputs=[tables_dropdown], outputs=[result_output, data_description, data_alerts, describe_cat, describe_num, tests_output, test_result_output])374 user_generate_result.click(user_results, inputs=[tables_dropdown, query_input], outputs=[query_tests, query_result])375 376if __name__ == "__main__":377 demo.launch(debug=True)378 379 380 