diff --git a/.gitignore b/.gitignore index 42d3722..7700ec7 100644 --- a/.gitignore +++ b/.gitignore @@ -4,8 +4,10 @@ __pycache__ Downloads # Environments -.env +.env +.sqlite3 .venv +.bin env/ venv/ ENV/ diff --git a/app/api/v1/llm.py b/app/api/v1/llm.py index e1ebc7f..902ff83 100644 --- a/app/api/v1/llm.py +++ b/app/api/v1/llm.py @@ -2,6 +2,11 @@ from fastapi import APIRouter, HTTPException from pydantic import BaseModel +from app.utils.llm.conversation_agent import ConversationAgent +from app.utils.llm.openai_model import OpenAIModel +from app.utils.llm.tools.reverse import create_your_own + + router = APIRouter() @@ -20,6 +25,14 @@ class PromptResponse(BaseModel): async def prompt(request_data: PromptRequest): try: reply_message = f"You said: {request_data.message}" + + tools = [create_your_own] # Define your tools + workroom_modal = OpenAIModel(tools=tools) # Initialize the OpenAI model + conversation_agent = ConversationAgent(openai_model=workroom_modal, tools=tools) + + reply_message = conversation_agent.convchain(request_data.message) + print(reply_message) + # Example of creating a pandas DataFrame (replace this with your actual data) data = {"column1": [1, 2, 3], "column2": [4, 5, 6]} df = pd.DataFrame(data) diff --git a/app/core/config.py b/app/core/config.py index d268cce..65efaa3 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -59,5 +59,5 @@ def get_config(env_state: str): # print(os.path.expanduser("~/.env")) # print(GlobalConfig().DATABASE_URL) # print(BaseConfig().ENV_STATE) -# print(get_config(BaseConfig().ENV_STATE)) +# print(get_config(BaseConfig())) config = get_config(BaseConfig().ENV_STATE) diff --git a/app/utils/__init__.py b/app/utils/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/app/utils/llm/__init__.py b/app/utils/llm/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/app/utils/llm/conversation_agent.py b/app/utils/llm/conversation_agent.py new file mode 100644 index 0000000..ddbf874 --- /dev/null +++ b/app/utils/llm/conversation_agent.py @@ -0,0 +1,38 @@ +from langchain.agents import AgentExecutor +from langchain.memory import ConversationBufferMemory +from langchain.prompts import ChatPromptTemplate, MessagesPlaceholder + +from app.utils.llm.sysprompt.conversational_agent import CONVERSATIONAL_AGENT_PROMPT + + + +class ConversationAgent: + def __init__(self, openai_model, tools): + self.openai_model = openai_model + self.memory = ConversationBufferMemory( + return_messages=True, memory_key="chat_history" + ) + self.prompt = self.create_prompt() + self.chain = self.openai_model.create_functional_chain(self.prompt) + self.qa = AgentExecutor( + agent=self.chain, tools=tools, verbose=True, memory=self.memory + ) + + def create_prompt(self): + return ChatPromptTemplate.from_messages( + [ + ( + "system", + CONVERSATIONAL_AGENT_PROMPT, + ), + MessagesPlaceholder(variable_name="chat_history"), + ("user", "{input}"), + MessagesPlaceholder(variable_name="agent_scratchpad"), + ] + ) + + def convchain(self, query): + if not query: + return + result = self.qa.invoke({"input": query}) + return result["output"] diff --git a/app/utils/llm/openai_model.py b/app/utils/llm/openai_model.py new file mode 100644 index 0000000..3b5d118 --- /dev/null +++ b/app/utils/llm/openai_model.py @@ -0,0 +1,25 @@ +from langchain.agents.format_scratchpad import format_to_openai_functions +from langchain.agents.output_parsers import OpenAIFunctionsAgentOutputParser +from langchain.chat_models import ChatOpenAI +from langchain.schema.runnable import RunnablePassthrough +from langchain.tools.render import format_tool_to_openai_function +from app.core.config import config + + +class OpenAIModel: + def __init__(self, tools): + + self.functions = [format_tool_to_openai_function(f) for f in tools] + self.chat_model = ChatOpenAI(openai_api_key=config.OPENAI_API_KEY,temperature=0).bind(functions=self.functions) + + def create_functional_chain(self, prompt): + return ( + RunnablePassthrough.assign( + agent_scratchpad=lambda x: format_to_openai_functions( + x["intermediate_steps"] + ) + ) + | prompt + | self.chat_model + | OpenAIFunctionsAgentOutputParser() + ) diff --git a/app/utils/llm/sysprompt/__init__.py b/app/utils/llm/sysprompt/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/app/utils/llm/sysprompt/conversational_agent.py b/app/utils/llm/sysprompt/conversational_agent.py new file mode 100644 index 0000000..837df36 --- /dev/null +++ b/app/utils/llm/sysprompt/conversational_agent.py @@ -0,0 +1,7 @@ +# sysprompt/conversational_agent.py + +CONVERSATIONAL_AGENT_PROMPT = "You are a helpful but sassy assistant. If you don't find anything, don't hallucinate it." +# CONVERSATIONAL_AGENT_PROMPT = ( +# "You are a helpful but sassy assistant. " +# "If you don't find anything, don't hallucinate it." +# ) \ No newline at end of file diff --git a/app/utils/llm/tools/__init__.py b/app/utils/llm/tools/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/app/utils/llm/tools/reverse.py b/app/utils/llm/tools/reverse.py new file mode 100644 index 0000000..db4cc2f --- /dev/null +++ b/app/utils/llm/tools/reverse.py @@ -0,0 +1,8 @@ +from langchain.tools import tool + + +@tool +def create_your_own(query: str) -> str: + """This function can do whatever you would like once you fill it in""" + print(type(query)) + return query[::-1] diff --git a/app/utils/llm/tools/web_search.py b/app/utils/llm/tools/web_search.py new file mode 100644 index 0000000..e0142c5 --- /dev/null +++ b/app/utils/llm/tools/web_search.py @@ -0,0 +1,52 @@ +import os + +# from langchain.chains import RetrievalQAWithSourcesChain +# from langchain.retrievers.web_research import WebResearchRetriever +# from langchain.tools import tool +# # from langchain.vectorstores import pgvector +# # from langchain_chroma import Chroma +# from langchain_community.utilities import GoogleSearchAPIWrapper +# from langchain_community.vectorstores import FAISS +# from langchain_openai import ChatOpenAI, OpenAIEmbeddings + +from app.core.config import config + +os.environ["GOOGLE_API_KEY"] = config.GOOGLE_API_KEY +os.environ["GOOGLE_CSE_ID"] = config.GOOGLE_CSE_ID +os.environ["OPENAI_API_KEY"] = config.OPENAI_API_KEY + + +# @tool +def search_web(query): + """ + Performs search on Google Shopping. + Args: + - query (str): search queries with specifications such as exact-match phrases ("nikola tesla"), logical OR operators (tesla OR edison), exclusion criteria (tesla -motors), wildcard usage (tesla "rock * roll"), range matching (tesla announcement 2015..2017), price searches (tesla deposit $1000), unit conversions (250 kph in mph), searches within page titles (intitle:"tesla vs edison"), URL searches (tesla announcements inurl:2016), text searches within document bodies (intext:"orbi vs google wifi"), filetype specific searches ("tesla announcements" filetype:pdf), related site searches (related:nytimes.com), proximity searches (tesla AROUND(3) edison), and chained operator combinations ("nikola tesla" intitle:"top 5..10 facts" -site:youtube.com inurl:2015). + + Returns: + - search results: string return. + """ + # vectorstore = Chroma( + # embedding_function=OpenAIEmbeddings(), persist_directory="./chroma_db_oai" + # ) + + # vectorstore = FAISS(embedding_function=OpenAIEmbeddings()) + + # # LLM + # llm = ChatOpenAI(openai_api_key=config.OPENAI_API_KEY, temperature=0) + + # # Search + # search = GoogleSearchAPIWrapper() + + # # Initialize + # web_research_retriever = WebResearchRetriever.from_llm( + # vectorstore=vectorstore, llm=llm, search=search + # ) + + # qa_chain = RetrievalQAWithSourcesChain.from_chain_type( + # llm, retriever=web_research_retriever + # ) + + # result = qa_chain({"question": query}) + result = True + return result diff --git a/app/utils/llm/tools/websearch.py b/app/utils/llm/tools/websearch.py deleted file mode 100644 index f785811..0000000 --- a/app/utils/llm/tools/websearch.py +++ /dev/null @@ -1,44 +0,0 @@ -import os -from langchain.chains import RetrievalQAWithSourcesChain -from langchain.retrievers.web_research import WebResearchRetriever -from langchain.tools import tool -from langchain_chroma import Chroma -from langchain_community.utilities import GoogleSearchAPIWrapper -from langchain_openai import ChatOpenAI, OpenAIEmbeddings -from app.core.config import config - -os.environ["GOOGLE_API_KEY"] = config.GOOGLE_API_KEY -os.environ["GOOGLE_CSE_ID"] = config.GOOGLE_CSE_ID - - -@tool -def search_web(query): - """ - Performs search on Google Shopping. - Args: - - query (str): search queries with specifications such as exact-match phrases ("nikola tesla"), logical OR operators (tesla OR edison), exclusion criteria (tesla -motors), wildcard usage (tesla "rock * roll"), range matching (tesla announcement 2015..2017), price searches (tesla deposit $1000), unit conversions (250 kph in mph), searches within page titles (intitle:"tesla vs edison"), URL searches (tesla announcements inurl:2016), text searches within document bodies (intext:"orbi vs google wifi"), filetype specific searches ("tesla announcements" filetype:pdf), related site searches (related:nytimes.com), proximity searches (tesla AROUND(3) edison), and chained operator combinations ("nikola tesla" intitle:"top 5..10 facts" -site:youtube.com inurl:2015). - - Returns: - - search results: string return. - """ - vectorstore = Chroma( - embedding_function=OpenAIEmbeddings(), persist_directory="./chroma_db_oai" - ) - - # LLM - llm = ChatOpenAI(temperature=0) - - # Search - search = GoogleSearchAPIWrapper() - - # Initialize - web_research_retriever = WebResearchRetriever.from_llm( - vectorstore=vectorstore, llm=llm, search=search - ) - - qa_chain = RetrievalQAWithSourcesChain.from_chain_type( - llm, retriever=web_research_retriever - ) - - result = qa_chain({"question": query}) - return result diff --git a/requirements.txt b/requirements.txt index 63b6c40..533d68e 100644 Binary files a/requirements.txt and b/requirements.txt differ