Team Ai
Apppublic

Benjamona97/sql-agent

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
app.py186 linesDownload Raw Back to root
1import re2from pathlib import Path3from typing import List4 5 6import chainlit as cl7from dotenv import load_dotenv8from langchain.pydantic_v1 import BaseModel, Field9from langchain.tools import StructuredTool10from langchain.indexes import SQLRecordManager, index11from langchain.schema import Document12from langchain.agents import initialize_agent, AgentExecutor13from langchain.text_splitter import RecursiveCharacterTextSplitter14from langchain.vectorstores.chroma import Chroma15from langchain_community.document_loaders import CSVLoader16from langchain_core.prompts import ChatPromptTemplate17from langchain_openai import ChatOpenAI, OpenAIEmbeddings18from openai import AsyncOpenAI19 20# from modules.database.database import PostgresDB21from modules.database.sqlitedatabase import Database22 23"""24Here we define some environment variables and the tools that the agent will use.25Along with some configuration for the app to start.26"""27load_dotenv()28 29chunk_size = 51230chunk_overlap = 5031 32embeddings_model = OpenAIEmbeddings()33openai_client = AsyncOpenAI()34 35CSV_STORAGE_PATH = "./data"36 37 38def remove_triple_backticks(text):39    # Use a regular expression to replace all occurrences of triple backticks with an empty string40    cleaned_text = re.sub(r"```", "", text)41    return cleaned_text42 43 44def process_pdfs(pdf_storage_path: str):45    csv_directory = Path(pdf_storage_path)46    docs = []  # type: List[Document]47    text_splitter = RecursiveCharacterTextSplitter(48        chunk_size=chunk_size, chunk_overlap=50)49 50    for csv_path in csv_directory.glob("*.csv"):51        loader = CSVLoader(file_path=str(csv_path))52        documents = loader.load()53        docs += text_splitter.split_documents(documents)54 55    documents_search = Chroma.from_documents(docs, embeddings_model)56 57    namespace = "chromadb/my_documents"58    record_manager = SQLRecordManager(59        namespace, db_url="sqlite:///record_manager_cache.sql"60    )61    record_manager.create_schema()62 63    index_result = index(64        docs,65        record_manager,66        documents_search,67        cleanup="incremental",68        source_id_key="source",69    )70 71    print(f"Indexing stats: {index_result}")72 73    return documents_search74 75 76doc_search = process_pdfs(CSV_STORAGE_PATH)77 78"""79Execute SQL query tool definition along schemas.80"""81 82 83def execute_sql(query: str) -> str:84    """85    Execute SQLite queries queries against the database. Delete all markdown code and backticks from the query.86    """87    db = Database("./db/mydatabase.db")88    db.connect()89 90    cleaned_query = remove_triple_backticks(query)91 92    results = db.execute_query(cleaned_query)93 94    return results + f"\nQuery used:\n```sql{cleaned_query}```"95 96 97class ExecuteSqlToolInput(BaseModel):98    query: str = Field(99        description="A SQLite query to be executed agains the database")100 101 102execute_sql_tool = StructuredTool(103    func=execute_sql,104    name="Execute SQL",105    description="useful for when you need to execute SQL queries against the database. Always use a clause LIMIT 10",106    args_schema=ExecuteSqlToolInput107)108 109"""110Research database tool definition along schemas.111"""112 113 114def research_database(user_request: str) -> str:115    """116    Searches for table definitions matching the user request117    """118    search_kwargs = {"k": 30}119 120    retriever = doc_search.as_retriever(search_kwargs=search_kwargs)121 122    def format_docs(docs):123        for i, doc in enumerate(docs):124            print(f"{i+1}. {doc.page_content}")125        return "\n\n".join([d.page_content for d in docs])126 127    results = retriever.invoke(user_request)128 129    return format_docs(results)130 131 132class ResearchDatabaseToolInput(BaseModel):133    user_request: str = Field(134        description="The user query to search against the table definitions for matches.")135 136 137research_database_tool = StructuredTool(138    func=research_database,139    name="Search db info",140    description="Search for database table definitions so you can have context for building SQL queries. The queries needs to be SQLite compatible.",141    args_schema=ResearchDatabaseToolInput142)143 144 145@cl.on_chat_start146def start():147    tools = [execute_sql_tool, research_database_tool]148 149    llm = ChatOpenAI(model="gpt-4", temperature=0, verbose=True)150 151    prompt = ChatPromptTemplate.from_template(152        """153            You are a SQLite world class data scientist, based on user query154            use your tools to do the job. Usually you would start by analyzing155            for possible SQL queries the user wants to build based on your knowledge base.156            Remember your tools are:157 158            - execute_sql (bring back the results as of running the query against the database)159            - research_database (search for table definitions so you can build a SQLite Query)160 161            Remember, you are building SQLite compatible queries. If you don't know the answer don't162            make anything up. Always ask for feedback. One last detail: always run the querys with LIMIT 10 and add163            the SQL query as markdown to the final answer so the user knows what SQL query was used for the job and164            can copy it for further use.165 166            REMEMBER TO GENERATE ALWAYS SQLITE COMPATIBLE QUERIES.167 168            User query: {input}169        """170    )171 172    agent = initialize_agent(tools=tools, prompt=prompt,173                             llm=llm, handle_parsing_errors=True)174 175    cl.user_session.set("agent", agent)176 177 178@cl.on_message179async def main(message: cl.Message):180    agent = cl.user_session.get("agent")  # type: AgentExecutor181    res = await agent.arun(182        message.content, callbacks=[cl.AsyncLangchainCallbackHandler()]183    )184 185    await cl.Message(content=res).send()186