ramhemanth580/NL_2_SQL_Data_Analysis_Chatbot
0
1import os2from dotenv import load_dotenv3 4import pandas as pd5import streamlit as st6from operator import itemgetter7#from langchain.chains.openai_tools import create_extraction_chain_pydantic8from langchain_core.pydantic_v1 import BaseModel, Field9#from langchain_openai import ChatOpenAI10from langchain.chains import LLMChain11from langchain_core.prompts import ChatPromptTemplate12 13import google.generativeai as genai14 15from langchain_google_genai import ChatGoogleGenerativeAI16 17 18from typing import List19 20load_dotenv()21genai.configure(api_key=os.environ["GOOGLE_API_KEY"])22llm = ChatGoogleGenerativeAI(model="gemini-pro",temperature=0,convert_system_message_to_human=True)23 24@st.cache_data25def get_table_details():26 # Read the CSV file into a DataFrame27 table_description = pd.read_csv("database_table_descriptions.csv")28 table_docs = []29 30 # Iterate over the DataFrame rows to create Document objects31 table_details = ""32 for index, row in table_description.iterrows():33 table_details = table_details + "Table Name:" + row['Table'] + "\n" + "Table Description:" + row['Description'] + "\n\n"34 35 return table_details36 37 38class Table(BaseModel):39 """Table in SQL database."""40 41 name: str = Field(description="Name of table in SQL database.")42 43table_details = get_table_details()44 45prompt2 = ChatPromptTemplate.from_template(46 """47 You are a helpful Data science assistant , Your objective is to analyze the following table descriptions and Return the names of ALL the SQL tables that MIGHT be relevant to the question: {question}48 \n\nRemember to include ALL POTENTIALLY RELEVANT tables, even if you're not sure that they're needed.and you should return the table names as a list49 for example question : which customers made the top 5 highest payments50 the desired answer should be ['customers','payments']51 \n\nThe tables descriptions are:52 Table Name:productlines53 Table Description:Stores information about the different product lines offered by the company, including a unique name, textual description, HTML description, and image. Categorizes products into different lines.54 55 Table Name:products56 Table Description:Contains details of each product sold by the company, including code, name, product line, scale, vendor, description, stock quantity, buy price, and MSRP. Linked to the productlines table.57 58 Table Name:offices59 Table Description:Holds data on the company's sales offices, including office code, city, phone number, address, state, country, postal code, and territory. Each office is uniquely identified by its office code.60 61 Table Name:employees62 Table Description:Stores information about employees, including number, last name, first name, job title, contact info, and office code. Links to offices and maps organizational structure through the reportsTo attribute.63 64 Table Name:customers65 Table Description:Captures data on customers, including customer number, name, contact details, address, assigned sales rep, and credit limit. Central to managing customer relationships and sales processes.66 67 Table Name:payments68 Table Description:Records payments made by customers, tracking the customer number, check number, payment date, and amount. Linked to the customers table for financial tracking and account management.69 70 Table Name:orders71 Table Description:Details each sales order placed by customers, including order number, dates, status, comments, and customer number. Linked to the customers table, tracking sales transactions.72 73 Table Name:orderdetails74 Table Description:Describes individual line items for each sales order, including order number, product code, quantity, price, and order line number. Links orders to products, detailing the items sold.75 76 """77)78 79from typing import List, Dict80import ast81 82# Assuming Table is a Pydantic model or similar83class Table:84 name: str85 86def get_tables(output: Dict) -> List[str]:87 # Extract the 'text' field from the output, which contains the list as a string88 text_output = output.get('text', '')89 90 try:91 # Safely evaluate the string representation of the list92 tables_list = ast.literal_eval(text_output)93 # Ensure that the result is indeed a list94 if isinstance(tables_list, list):95 # Extract the table names if 'tables_list' is a list of Table objects96 # If it's already a list of strings, you can return it directly97 return [table.name if isinstance(table, Table) else table for table in tables_list]98 except (ValueError, SyntaxError):99 # Handle the case where the text output is not a valid list representation100 return []101 102table_chain = {"question": itemgetter("question")} | LLMChain(llm=llm, prompt=prompt2) | get_tables103 104 105# table_names = "\n".join(db.get_usable_table_names())106# table_details = get_table_details()107# table_details_prompt = f"""Return the names of ALL the SQL tables that MIGHT be relevant to the user question. \108# The tables are:109 110# {table_details}111 112# Remember to include ALL POTENTIALLY RELEVANT tables, even if you're not sure that they're needed."""113 114# table_chain = {"input": itemgetter("question")} | create_extraction_chain_pydantic(Table, llm, system_message=table_details_prompt) | get_tables