rolzy/sql_chatbot
0
1import os2import dotenv3import gradio as gr4 5# LLMs6import openai7from langchain.chat_models import AzureChatOpenAI8from langchain.llms import AzureOpenAI9from langchain.schema import AIMessage, HumanMessage10from langchain_experimental.sql import SQLDatabaseChain11from langchain.agents import create_sql_agent12from langchain.agents.agent_toolkits import SQLDatabaseToolkit13from langchain.agents.agent_types import AgentType14from langchain.agents.mrkl.prompt import FORMAT_INSTRUCTIONS15 16# Databases17import adal18import struct19from sqlalchemy.engine import URL20from sqlalchemy import create_engine, event21from langchain import SQLDatabase22 23# Load environment variables24dotenv.load_dotenv()25 26# Set OpenAI API settings from the .env file27openai.api_type = "azure"28openai.api_base = os.getenv("OPENAI_API_BASE")29openai.api_key = os.getenv("OPENAI_API_KEY")30openai.api_version = os.getenv('OPENAI_API_VERSION')31 32def get_token():33 tenantId = os.getenv("DB_TENANT_ID")34 clientId = os.getenv("DB_CLIENT_ID")35 clientSecret = os.getenv("DB_CLIENT_SECRET")36 37 authorityHostUrl = "https://login.microsoftonline.com"38 authorityUrl = authorityHostUrl + "/" + tenantId39 context = adal.AuthenticationContext(authorityUrl, api_version=None)40 token = context.acquire_token_with_client_credentials("https://database.windows.net/", clientId, clientSecret)41 42 tokenb = bytes(token["accessToken"], "UTF-8")43 exptoken = b''44 for i in tokenb:45 exptoken += bytes({i})46 exptoken += bytes(1)47 return struct.pack("=i", len(exptoken)) + exptoken48 49def get_conn_url():50 server = "sql-ae-dsdev-dna-hack01-pnpdebzuohv7y.database.windows.net"51 database = 'dna-hack-db01'52 return f"mssql+pyodbc://@{server}/{database}?driver=ODBC+Driver+17+for+SQL+Server"53 54def get_database():55 # connection_string = "mssql+pyodbc://@my-server.database.windows.net/myDb?driver=ODBC+Driver+17+for+SQL+Server"56 conn_url = get_conn_url()57 print(conn_url)58 engine = create_engine(conn_url)59 60 @event.listens_for(engine, "do_connect")61 def provide_token(dialect, conn_rec, cargs, cparams):62 # remove the "Trusted_Connection" parameter that SQLAlchemy adds63 cargs[0] = cargs[0].replace(";Trusted_Connection=Yes", "")64 65 # create token credential66 token_struct = get_token()67 68 # apply it to keyword arguments69 SQL_COPT_SS_ACCESS_TOKEN = 125670 cparams["attrs_before"] = {SQL_COPT_SS_ACCESS_TOKEN: token_struct}71 72 return SQLDatabase(engine = engine, schema="[dntmatrix]")73 74def get_llm():75 return AzureChatOpenAI(temperature=1.0, 76 model_name='gpt-4',77 deployment_name='gpt-4',78 model_kwargs={79 "engine": "gpt-4",80 "api_key": openai.api_key,81 "api_base": openai.api_base,82 "api_type": openai.api_type,83 "api_version": openai.api_version84 }85 )86 87def get_sql_chain():88 db = get_database()89 llm = get_llm()90 return SQLDatabaseChain.from_llm(llm, db, verbose=True)91 92def get_sql_agent():93 db = get_database()94 llm = get_llm()95 format_instruction = """96 97 """ + FORMAT_INSTRUCTIONS98 99 toolkit = SQLDatabaseToolkit(db=db, llm=llm)100 return create_sql_agent(101 llm=llm,102 toolkit=toolkit,103 verbose=True,104 agent_type=AgentType.ZERO_SHOT_REACT_DESCRIPTION,105 format_instructions=format_instruction106 )107 108sql_agent = get_sql_agent()109 110def predict(message, history):111 gpt_response = sql_agent.run(message)112 return gpt_response113 114with gr.Blocks() as demo:115 gr.ChatInterface(predict)116 117if __name__ == "__main__":118 demo.launch(server_name="0.0.0.0")119 