Update agent.py

This commit is contained in:
2026-06-02 07:15:38 +00:00
parent 59ce21aacf
commit 883afee0c3
+69 -64
View File
@@ -1,18 +1,16 @@
"""Main script for the RAG agent with ChromaDB and Tavily. """Main RAG agent implementation.
The script: The agent can answer questions using either the local ChromaDB knowledge base
1. Loads or creates the Chroma vector store. or live web search via Tavily. The decision of which tool to use is made by
2. Loads documents from the `documents/` folder. the LLM itself based on the prompt.
3. Sets up the LangChain agent with two tools: `search_local_kb` and `web_search`.
4. Runs a simple CLI loop.
""" """
import os import os
from pathlib import Path from pathlib import Path
from langchain_ollama import ChatOllama from langchain_ollama import ChatOllama
from langchain.agents import initialize_agent, AgentType from langchain.agents import AgentExecutor, create_openai_tools_agent
from langchain.tools import Tool from langchain.prompts import ChatPromptTemplate, SystemMessagePromptTemplate, HumanMessagePromptTemplate
from vectorstore import create_vectorstore, load_documents from vectorstore import create_vectorstore, load_documents
from rag_tools import search_local_kb, web_search from rag_tools import search_local_kb, web_search
@@ -20,80 +18,87 @@ from rag_tools import search_local_kb, web_search
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Configuration # Configuration
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
CHROMA_DIR = "./chroma_db" VECTORSTORE_DIR = Path("./chroma_db")
DOCS_DIR = "./documents" DOCUMENTS_DIR = Path("./documents")
MODEL = "llama3"
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Helper: load or create vector store # Initialise vector store and retriever
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
vectorstore = create_vectorstore(persist_directory=CHROMA_DIR) vectorstore = create_vectorstore(str(VECTORSTORE_DIR))
# Load documents on first run this is idempotent
# Load documents we always load; Chroma will deduplicate by ID if same content if not any(VECTORSTORE_DIR.iterdir()):
print("Loading documents into ChromaDB (if not already present)...") print("Loading documents into ChromaDB")
load_documents(DOCS_DIR, vectorstore) load_documents(str(DOCUMENTS_DIR), vectorstore)
print("Documents loaded.") print("Documents loaded.")
# --------------------------------------------------------------------------- # Global retriever for tool access
# Define tools pass the vectorstore to the local search tool vectorstore_retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
# ---------------------------------------------------------------------------
# We wrap the tool functions to include the vectorstore argument
def local_kb_tool(query: str, top_k: int = 3): # ---------------------------------------------------------------------------
return search_local_kb(query=query, top_k=top_k, vectorstore=vectorstore) # LLM and prompt
# ---------------------------------------------------------------------------
llm = ChatOllama(model="llama3")
# Create LangChain Tool objects system_prompt = """You are an AI assistant that can answer questions using two sources:
local_tool = Tool(
name="search_local_kb", 1. A local knowledge base (ChromaDB). Use the tool ``search_local_kb`` when the
func=local_kb_tool, answer can be found in the documents.
description="Semantic search in the local knowledge base. Use when the answer is in the local documents.", 2. Live web search (Tavily). Use the tool ``web_search`` when the answer requires
) uptodate information.
web_tool = Tool(
name="web_search", After retrieving the information, answer the user question and explicitly
func=web_search, state the source you used: either ``chromadb`` or ``tavily``.
description="Search the web using Tavily. Use for uptodate facts or news.",
) If you are unsure, ask for clarification. Do not provide fabricated data.
"""
prompt = ChatPromptTemplate.from_messages([
SystemMessagePromptTemplate.from_template(system_prompt),
HumanMessagePromptTemplate.from_template("{input}")
])
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Agent setup # Agent setup
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
llm = ChatOllama(model=MODEL, temperature=0.0) # Tools are automatically discovered via the @tool decorator in rag_tools.py
tools = [search_local_kb, web_search]
system_prompt = ( agent = create_openai_tools_agent(llm=llm, tools=tools, prompt=prompt)
"You are an assistant that answers user questions. " agent_executor = AgentExecutor.from_agent_and_tools(agent=agent, tools=tools, verbose=True)
"If the answer is likely to be in the local knowledge base, use the tool "
"search_local_kb. If the answer requires uptodate information, use the "
"web_search tool. After retrieving information, provide the answer and "
"state the source: either 'chromadb' or 'tavily'."
)
agent = initialize_agent(
tools=[local_tool, web_tool],
llm=llm,
agent=AgentType.CHAT_ZERO_SHOT_REACT_DESCRIPTION,
verbose=True,
prefix=system_prompt,
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# CLI loop # Public API
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
print("\nRAG Agent ready. Type your question (or 'exit' to quit).\n")
def answer_query(query: str) -> str:
"""Return the agent's answer for *query*.
Parameters
----------
query: str
The user's question.
Returns
-------
str
The agent's response.
"""
result = agent_executor.invoke({"input": query})
return result["output"]
# ---------------------------------------------------------------------------
# CLI entry point
# ---------------------------------------------------------------------------
if __name__ == "__main__":
print("RAG Agent ready. Type 'exit' to quit.")
while True: while True:
try: try:
query = input("Query: ") user_input = input("\nQuery: ")
except (KeyboardInterrupt, EOFError): except (KeyboardInterrupt, EOFError):
print("\nExiting.") print("\nExiting.")
break break
if query.strip().lower() in {"exit", "quit", "q"}: if user_input.lower() in {"exit", "quit"}:
print("Exiting.") print("Goodbye!")
break break
if not query.strip(): response = answer_query(user_input)
continue print("\nAnswer:\n", response)
# Run the agent
try:
result = agent.run(query)
print(f"\nAnswer:\n{result}\n")
except Exception as e:
print(f"Error: {e}")
continue