SoumyaJ/DatabaseToolkitWithRAG
0
1import streamlit as st 2from langchain.agents import create_sql_agent,create_react_agent3from langchain.agents.agent_toolkits import SQLDatabaseToolkit4from langchain.agents.agent_types import AgentType5from langchain_groq import ChatGroq6from langchain_core.prompts import ChatPromptTemplate7from langchain.sql_database import SQLDatabase8from sqlalchemy import create_engine9from langchain.text_splitter import RecursiveCharacterTextSplitter10from langchain_huggingface import HuggingFaceEmbeddings11from langchain_community.vectorstores import FAISS12from langchain.chains.combine_documents import create_stuff_documents_chain13from langchain.chains.retrieval import create_retrieval_chain14from langchain_core.output_parsers import StrOutputParser15from sqlalchemy.orm import sessionmaker16from sqlalchemy import text17import sqlite318from dotenv import load_dotenv19from pathlib import Path20from PyPDF2 import PdfReader21import os22import re23 24load_dotenv()25os.environ['GROQ_API_KEY'] = os.getenv("GROQ_API_KEY") 26os.environ['HF_TOKEN'] = os.getenv("HF_TOKEN")27 28st.set_page_config("Langchain interaction with DB")29st.title("Document QnA with DB interaction")30 31llm = ChatGroq(model="llama3-8b-8192", api_key= os.environ['GROQ_API_KEY'])32 33embeddings = HuggingFaceEmbeddings(model_name = "all-MiniLM-L6-v2")34 35duration_pattern = re.compile(r"(\d+)\s*(min[s]?|minute[s]?)")36 37st.session_state.user_prompt = ""38st.session_state.summary = ""39 40pdf_prompt_template = ChatPromptTemplate.from_template("""41Answer the following question from the provided context only. 42Please provide the most accurate response based on the question43<context>44{context}45</context> 46Question : {input}47""")48 49def get_pdf_text(pdf_docs):50 text=""51 for pdf in pdf_docs:52 pdf_reader= PdfReader(pdf)53 for page in pdf_reader.pages:54 text+= page.extract_text()55 return text56 57def create_vector_embeddings(pdfText):58 if "vectors" not in st.session_state: 59 st.session_state.docs = get_pdf_text(pdfText)60 st.session_state.splitter = RecursiveCharacterTextSplitter(chunk_size=1200,chunk_overlap=400)61 st.session_state.final_docs = st.session_state.splitter.split_text(st.session_state.docs) 62 st.session_state.vectors = FAISS.from_texts(st.session_state.final_docs, embeddings)63 64def configure():65 dbfilepath = (Path(__file__).parent /"programme.db").absolute()66 creator = lambda: sqlite3.connect(f"file:{dbfilepath}",uri= True, check_same_thread=False)67 return create_engine("sqlite:///", creator= creator)68 69engine = configure()70db = SQLDatabase(engine)71#ChatGroq(model="gemma2-9b-it"72sql_toolkit = SQLDatabaseToolkit(db = db, llm = llm , api_key= os.environ['GROQ_API_KEY'])73sql_toolkit.get_tools()74 75prefilled_prompt = ""76 77# if "uploaded_text" in st.session_state:78# for m in st.session_state.uploaded_text:79# st.error(m)80# if 'PACKAGE' in st.session_state.uploaded_text:81# prefilled_prompt = "get the entire programme details linked to the package"82# else:83# prefilled_prompt = "get the entire programme details linked to the document"84 85# query=st.text_input("ask question here", value = prefilled_prompt)86 87def clear_database(): 88 89 connection = engine.raw_connection()90 try:91 # Create a cursor from the raw connection92 cursor = connection.cursor()93 94 # List of tables to clear95 tables = ["programme", "episode"]96 97 # Execute DELETE commands for each table98 for table in tables:99 cursor.execute(f"DELETE FROM {table}")100 101 # Commit the changes to the database102 connection.commit()103 finally:104 # Ensure the connection is closed properly105 connection.close()106 107 108def process_sql_script(sql_script):109 # Define the keyword to check110 keyword = 'PACKAGE'111 112 # Split the script into lines113 lines = sql_script.strip().split(';')114 115 programme_line = lines[0]116 if keyword not in programme_line: 117 filtered_script = "\n".join([lines[0]])118 else: 119 filtered_script = "\n".join(lines)120 121 return filtered_script122 123import re124 125def convert_to_hms(duration): 126 hour_minute_match = re.match(r'(?:(\d+)\s*hour[s]?)?\s*(\d+)\s*min[s]?', duration.lower())127 128 if hour_minute_match:129 hours = int(hour_minute_match.group(1) or 0)130 minutes = int(hour_minute_match.group(2) or 0)131 else:132 return duration133 134 total_seconds = (hours * 60 * 60) + (minutes * 60)135 hh = total_seconds // 3600136 mm = (total_seconds % 3600) // 60137 ss = total_seconds % 60138 139 return f"{hh:02}:{mm:02}:{ss:02}"140 141def handleDurationForEachScript(scripts):142 filtered_data = ""143 # for script in scripts.split(";"):144 # # Find all matches for durations like '60 minutes' or '60 mins'145 # matches = duration_pattern.findall(script)146 147 # for match in matches:148 # duration = f"{match[0]} {match[1]}" # e.g., '60 mins' or '60 minutes'149 # converted_duration = convert_to_hms(duration) # Convert to hh:mm:ss150 # script = script.replace(duration, converted_duration).replace('utes','') # Replace in script151 # if ('episode' not in filtered_data) & ('programme' not in filtered_data):152 # filtered_data = filtered_data + script153 pattern = r"'(\d+\s*(?:mins|minutes))'"154 for script in scripts.split(";"):155 match = re.search(pattern, script) 156 if match:157 duration = match.group(1) 158 converted_duration = convert_to_hms(duration) # Convert to hh:mm:ss159 script = script.replace(duration, converted_duration).replace('utes','') # Replace in script160 if ('episode' not in filtered_data) & ('programme' not in filtered_data):161 filtered_data = filtered_data + script 162 163 return filtered_data164 165def parse_insert_statement(insert_statement):166 # Extract the table name167 table_match = re.search(r'INSERT INTO (\w+)', insert_statement)168 if not table_match:169 return None, None, None170 171 table = table_match.group(1)172 173 # Extract columns and values174 columns_match = re.search(r'\((.*?)\)', insert_statement, re.DOTALL)175 values_match = re.search(r'VALUES\s*\((.*?)\)', insert_statement, re.DOTALL)176 177 if not columns_match or not values_match:178 return None, None, None179 180 columns = columns_match.group(1).replace('"', '').replace('\n', ' ').strip()181 values = values_match.group(1).replace("'", "").replace('\n', ' ').strip()182 183 return table, columns, values184 185def build_data_from_sql(programme_sql, episode_sql=None):186 data = {187 'Table': [],188 'Columns': [],189 'Values': []190 }191 192 # Parse the programme insert statement193 programme_table, programme_columns, programme_values = parse_insert_statement(programme_sql)194 195 if programme_table and programme_columns and programme_values:196 data['Table'].append(programme_table.capitalize()) 197 data['Columns'].append(programme_columns)198 data['Values'].append(programme_values)199 200 # Parse the episode insert statement, if it exists201 if episode_sql:202 episode_table, episode_columns, episode_values = parse_insert_statement(episode_sql)203 204 if episode_table and episode_columns and episode_values:205 data['Table'].append(episode_table.capitalize()) 206 data['Columns'].append(episode_columns)207 data['Values'].append(episode_values)208 209 return data210 211with st.sidebar:212 st.title("Menu:")213 #if "uploaded_text" not in st.session_state:214 st.session_state.uploaded_text = st.file_uploader("Upload your Files and Click on the Submit & Process Button", accept_multiple_files=True) 215 if st.button("Click To Process File"):216 with st.spinner("Processing..."):217 create_vector_embeddings(st.session_state.uploaded_text)218 st.write("Vector Database is ready") 219 220 # if "uploaded_text" in st.session_state and st.session_state.uploaded_text is not None: 221 # uploaded_file_names = [file.name for file in st.session_state.uploaded_text] 222 # if any('PACKAGE' in file_name.upper() for file_name in uploaded_file_names):223 # prefilled_prompt = "get the entire programme details linked to the package"224 # else: 225 # prefilled_prompt = "get the entire programme details linked to the document"226 227query=st.text_input("ask question here")228 229if query and "vectors" in st.session_state:230 st.session_state.user_prompt = query231 document_chain = create_stuff_documents_chain(llm=llm, prompt= pdf_prompt_template)232 retriever = st.session_state.vectors.as_retriever()233 retrieval_chain=create_retrieval_chain(retriever,document_chain)234 response = retrieval_chain.invoke({"input": st.session_state.user_prompt})235 #st.write(response)236 if response:237 st.session_state.summary = response['answer']238 st.write(response['answer'])239 240prompt=ChatPromptTemplate.from_messages(241 [242 ("system",243 """244 You are a SQL expert. Your task is to generate SQL INSERT scripts based on the provided context.245 246 1. Generate an `INSERT` statement for the `programme` table using the following values:247 - `ProgrammeTitle`248 - `ProgrammeType`249 - `Genre`250 - `SubGenre`251 - `Language`252 - `Duration` 253 Example:254 255 2562. After generating the `programme` statement, check the `ProgrammeTitle`:257 - If the `ProgrammeTitle` contains the keyword `PACKAGE`, generate an additional `INSERT` statement for the `episode` table.258 - If the `ProgrammeTitle` does **not** contain the keyword `PACKAGE`, **do not** generate an `INSERT` statement for the `episode` table.259 2603. The `episode` INSERT statement should look like this if the condition is met. EpisodeNumber is always 1 and `EpisodeTitle` should take same data from `ProgrammeTitle`.261 262 2634. Include only the SQL insert script(s) as final answer, **donot** include any additional details and notes.Return only the necessary SQL INSERT script(s) based on the current input. Ensure that no `episode` INSERT statement is included if the `ProgrammeTitle` does not contain `'PACKAGE'`.264 265Your output should strictly follow these conditions. Output **only** the final answer without producing any intermediate actions.266 267 """268 ),269 ("user","{question}\ ai: ")270 ])271 272agent=create_sql_agent(llm=llm,toolkit=sql_toolkit,agent_type=AgentType.ZERO_SHOT_REACT_DESCRIPTION,verbose=True,max_execution_time=100,max_iterations=1000, handle_parsing_errors=True)273 274if st.button("Generate Scripts",type="primary"): 275 try: 276 if st.session_state.summary is not None: 277 response=agent.run(prompt.format_prompt(question=st.session_state.summary))278 #with st.expander("Expand here to view scripts"): 279 if "INSERT" in response: 280 final_response = process_sql_script(response) 281 final_response_new = handleDurationForEachScript(final_response) 282 episode_sql = ""283 splitted_data = []284 285 if "uploaded_text" in st.session_state and st.session_state.uploaded_text is not None: 286 uploaded_file_names = [file.name for file in st.session_state.uploaded_text] 287 if any('PACKAGE' in file_name.upper() for file_name in uploaded_file_names):288 if ";" in final_response_new: 289 splitted_data = [stmt.strip() for stmt in final_response_new.strip().split(';') if stmt.strip()]290 elif "\n" in final_response_new:291 splitted_data = [stmt.strip() for stmt in final_response_new.strip().split('\n') if stmt.strip()]292 elif "," in final_response_new:293 splitted_data = [stmt.strip() for stmt in final_response_new.strip().split(',') if stmt.strip()]294 else:295 if final_response_new is list:296 splitted_data = final_response_new297 else:298 splitted_data.append(final_response_new)299 300 print(splitted_data)301 if len(splitted_data) > 0:302 programme_sql = splitted_data[0] + ';' # Re-add semicolon to the programme SQL statement303 print(f"prog{programme_sql}")304 if len(splitted_data) > 1:305 episode_sql = splitted_data[1] 306 #print(f"eps{episode_sql}")307 308 309 data = build_data_from_sql(programme_sql, episode_sql)310 st.write("### Script Summary")311 st.table(data)312 313 st.write("### Full SQL Scripts")314 315 with st.expander("Insert Scripts"):316 st.code(programme_sql, language='sql') 317 st.code(episode_sql, language='sql')318 319 #if episode_sql:320 #with st.expander("Episode Insert Script"):321 #st.code(episode_sql, language='sql')322 #st.code(final_response_new, language = 'sql')323 clear_database()324 325 #st.write(response)326 except Exception as e:327 st.error(f"Parsing error from LLM.Retry again !!! \n : {str(e)}") 328 329 330 331# data = {332# 'Table': ['Programme', 'Episode'],333# 'Columns': ['ProgrammeTitle, ProgrammeType, ...', 'EpisodeTitle, EpisodeNumber, ...'],334# 'Values': ['CHAMSARANG PACKAGE, Series, ...', 'CHAMSARANG PACKAGE, 1, ...']335# }336 337# # Display summary table338# st.write("### Script Summary")339# st.table(data)340 341# # Display expandable sections for each script342# st.write("### Full SQL Scripts")343# with st.expander("Programme Insert Script"):344# st.code("INSERT INTO programme ...")345 346# with st.expander("Episode Insert Script"):347# st.code("INSERT INTO episode ...")