CoolFace
Apppublic

SoumyaJ/DatabaseToolkitWithRAG

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py347 linesDownload Raw Back to root
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 ...")