ahmad-sobeh/fastapi-sqlite-crud
0
1#from orc_conn import db_connect
2#from find_schema_relationship import find_relations_from_file
3#db_connect
4
5#find_relations_from_file('my_columns.csv', 'suggested_relations.csv')
6import pandas as pd
7import nltk
8from nltk.stem import WordNetLemmatizer
9
10
11# Download the required dictionary for lemmatization
12nltk.download('wordnet')
13lemmatizer = WordNetLemmatizer()
14
15def before_nl2sql(question: str, csv_path: str = "my_columns.csv"):
16 df = pd.read_csv(csv_path)
17
18 # 1. Exclusion filter
19 exclude_keywords = ['test', 'temp', 'log', 'backup', 'dummy', 'tmp']
20 pattern = '|'.join(exclude_keywords)
21 df = df[~df['table_name'].str.contains(pattern, case=False, na=False)]
22
23 # 2. Lemmatize the user's question words (e.g., "customers" -> "customer")
24 # We also keep the original words to be safe
25 raw_words = question.lower().split()
26 normalized_words = set(raw_words + [lemmatizer.lemmatize(w) for w in raw_words])
27
28 def row_matches(row):
29 # Normalize the table and column names in the CSV as well
30 t_name = str(row['table_name']).lower()
31 c_name = str(row['column_name']).lower()
32
33 # Check if any normalized word from the user matches the schema
34 return (
35 any(w in t_name for w in normalized_words) or
36 any(w in c_name for w in normalized_words) or
37 # Also check if the table name lemmatized matches (e.g. schema has 'customers')
38 lemmatizer.lemmatize(t_name) in normalized_words
39 )
40
41 mask = df.apply(row_matches, axis=1)
42 relevant_df = df[mask] if mask.any() else df.iloc[:0]
43
44 formatted_schema = ""
45 export_data = []
46 for table, group in relevant_df.groupby('table_name'):
47 cols_str = ", ".join(group['column_name'].tolist())
48 formatted_schema += f"Table: {table} ({cols_str})\n"
49 export_data.append({"table_name": table, "columns": cols_str})
50
51 pd.DataFrame(export_data).to_csv("output.csv", index=False)
52 print(formatted_schema)
53 return formatted_schema
54
55before_nl2sql("get count for all customers")
56
57 