Benjamona97/sql-agent
0
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 