dvwn/nl2sql-api
0
1# Path: src/scripts/interactive_mode.py2# Interactive mode: Allows user to manually type questions and see the agent's response3import csv4import json5import pandas as pd6from pathlib import Path7from tabulate import tabulate8from src.database.db_manager import get_db_connection9from src.nl2sql.sql_agent import nl2sql_agent10from src.nl2sql.hf_engine import get_models11 12TEST_CASES_PATH = Path("src/scripts/test_cases.json")13 14def get_query_data(sql_query: str) -> pd.DataFrame:15 """16 Executes a SQL query and returns the results as a DataFrame.17 """18 if not sql_query or sql_query == "N/A":19 return pd.DataFrame()20 21 connection = get_db_connection()22 if not connection:23 return pd.DataFrame()24 25 try:26 df = pd.read_sql_query(sql_query, connection)27 return df28 except Exception as e:29 print(f"Error executing SQL query: {e}")30 return pd.DataFrame()31 finally:32 connection.close()33 34def verify_data(df_gold: pd.DataFrame, df_generated: pd.DataFrame) -> bool:35 """36 Guardrail check:37 Verifies if the generated DataFrame matches the expected gold DataFrame.38 """39 if df_gold.empty and df_generated.empty:40 return False # Both empty: To catch this as a potential issue41 42 if len(df_gold) != len(df_generated):43 return False44 45 try:46 gold_val = df_gold.fillna("").astype(str).values.tolist()47 gen_val = df_generated.fillna("").astype(str).values.tolist()48 49 # Strip whitespace and convert to tuples (sortable)50 gold_tuples = [tuple(val.strip() for val in row) for row in gold_val]51 gen_tuples = [tuple(val.strip() for val in row) for row in gen_val]52 53 return sorted(gold_tuples) == sorted(gen_tuples)54 except Exception as e:55 print(f"Error during data verification: {e}")56 return False57 58 59def run_interactiveMode(model_id: str):60 """ 61 Automates the interactive Question Answering mode.62 Runs predefined questions through the agent and logs the textual NL response.63 """64 print("\n========= Interactive NL2SQL Mode =========")65 print(f"Running Interactive Question Answering Evaluation on Model: {model_id}")66 #print("Type 'exit' or 'q' to return to the main menu.\n")67 68 if not TEST_CASES_PATH.exists():69 print(f"Error: Could not find test cases at {TEST_CASES_PATH}")70 return71 72 with TEST_CASES_PATH.open("r", encoding="utf-8") as handle:73 test_cases = json.load(handle)74 75 questAns_results = []76 77 for case in test_cases:78 case_id = case.get("id")79 question = case.get("question")80 gold_sql = case.get("gold_sql")81 print(f"\n\nTesting ID {case_id}: {question[:40]}...")82 83 response = nl2sql_agent(user_question=question, model_id=model_id)84 85 # Extract metadata86 status = response.get('status')87 nl_answer = response.get('nl_response', 'N/A')88 sql_query = response.get('query', 'N/A')89 error_msg = response.get('error', '')90 attempts = response.get('attempts', 0)91 92 # Data cross-check93 df_gold = get_query_data(gold_sql)94 df_generated = get_query_data(sql_query)95 96 # Verify accuracy97 is_data_accurate = verify_data(df_gold, df_generated)98 99 questAns_results.append({100 "id": case_id,101 "model_id": model_id,102 "question": question,103 "status": status,104 "data_returned_correct": is_data_accurate,105 "attempts": attempts,106 "nl_response": nl_answer,107 "sql_generated": sql_query,108 "error": error_msg109 })110 111 # Save to CSV for human-readable and easy comparison112 safe_model_name = model_id.replace("/", "_").replace(":", "_").replace(" ", "_")113 output_csv = Path(f"Q&A_report_{safe_model_name}.csv")114 115 keys = questAns_results[0].keys()116 with output_csv.open("w", newline='', encoding="utf-8") as f:117 dict_writer = csv.DictWriter(f, fieldnames=keys)118 dict_writer.writeheader()119 dict_writer.writerows(questAns_results)120 121 print(f"\nInteractive evaluation completed. Results saved to: {output_csv}")122 123if __name__ == "__main__":124 from dotenv import load_dotenv, find_dotenv125 load_dotenv(find_dotenv())126 127 models_to_test = get_models()128 for model in models_to_test:129 run_interactiveMode(model)