Spaces:
Sleeping
Sleeping
| import os, streamlit as st | |
| from llama_index import GPTVectorStoreIndex, SimpleDirectoryReader, LLMPredictor, PromptHelper, ServiceContext | |
| from llama_index import SQLDatabase,ServiceContext | |
| from llama_index.llms import OpenAI | |
| from sqlalchemy import select, create_engine, MetaData, Table,inspect | |
| from llama_index.indices.struct_store.sql_query import NLSQLTableQueryEngine | |
| from llama_index.prompts.prompts import SimpleInputPrompt | |
| import openai | |
| openai_api_key = st.sidebar.text_input( | |
| label="#### Your OpenAI API key π", | |
| placeholder="Paste your openAI API key, sk-", | |
| type="password") | |
| db_user = st.sidebar.text_input( | |
| label="#### Enter mysql user π", | |
| placeholder="Enter mysql username", | |
| type="default") | |
| db_password = st.sidebar.text_input( | |
| label="#### Enter mysql password π", | |
| placeholder="Enter mysql password", | |
| type="password") | |
| db_host = st.sidebar.text_input( | |
| label="#### Enter mysql server address π", | |
| placeholder="Enter mysql server address", | |
| type="default") | |
| db_port = st.sidebar.text_input( | |
| label="#### Enter mysql server port π", | |
| placeholder="Enter mysql server port", | |
| type="default") | |
| db_name = st.sidebar.text_input( | |
| label="#### Enter mysql db name π", | |
| placeholder="Enter mysql db name", | |
| type="default") | |
| def get_response(question: str | list[str],synthesize_response: bool = True): | |
| SYSTEM_PROMPT = """You are an AI assistant that answers questions in a friendly manner, based on the given source documents. Here are some rules you always follow: | |
| - Generate human readable output, avoid creating output with gibberish text. | |
| - Generate only the requested output, don't include any other language before or after the requested output. | |
| - Never say thank you, that you are happy to help, that you are an AI agent, etc. Just answer directly. | |
| - Generate professional language typically used in business documents in North America. | |
| - Never generate offensive or foul language. | |
| """ | |
| query_wrapper_prompt = "[INST]<<SYS>>\n" + SYSTEM_PROMPT + f"<</SYS>>\n\n {question}[/INST]" | |
| os.environ['OPENAI_API_KEY'] = openai_api_key | |
| openai.api_key = os.environ.get('OPENAI_API_KEY') | |
| llm = OpenAI(temperature=0.5, model="gpt-3.5-turbo-16k") | |
| connection_uri = f"mysql+pymysql://{db_user}:{db_password}@{db_host}:{db_port}/{db_name}" | |
| service_context = ServiceContext.from_defaults(llm=llm) | |
| engine = create_engine(connection_uri) | |
| sql_database = SQLDatabase(engine) | |
| inspector = inspect(engine) | |
| table_names = inspector.get_table_names() | |
| response_template = """ | |
| ## Question | |
| {question} | |
| ## Answer | |
| ``` | |
| {response} | |
| ``` | |
| ## Generated SQL Query | |
| ``` | |
| {sql} | |
| ``` | |
| """ | |
| query_engine = NLSQLTableQueryEngine( | |
| sql_database=sql_database, | |
| tables=table_names, | |
| synthesize_response=synthesize_response, | |
| service_context=service_context, | |
| ) | |
| try: | |
| response = query_engine.query(query_wrapper_prompt) | |
| response_md = str(response) | |
| sql = response.metadata["sql_query"] | |
| st.success(""" | |
| ## Question | |
| {} | |
| ## Answer | |
| {} | |
| ## Generated SQL Query | |
| {} | |
| """.format(question,response,sql)) | |
| except Exception as ex: | |
| st.error(f"ERROR: {str(ex)}") | |
| # Define a simple Streamlit app | |
| st.title("DbChat") | |
| query = st.text_input("What would you like to ask to db?", "") | |
| # If the 'Submit' button is clicked | |
| if st.button("Submit"): | |
| if not query.strip(): | |
| st.error(f"Please provide the search query.") | |
| else: | |
| try: | |
| if len(openai_api_key) > 0: | |
| get_response(query) | |
| else: | |
| st.error(f"Enter a valid openai key") | |
| except Exception as e: | |
| st.error(f"An error occurred: {e}") |