CoolFace
Apppublic

rolzy/sql_chatbot

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
main.py119 linesDownload Raw Back to root
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