Update agent.py
This commit is contained in:
@@ -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
|
||||||
|
if not any(VECTORSTORE_DIR.iterdir()):
|
||||||
|
print("Loading documents into ChromaDB…")
|
||||||
|
load_documents(str(DOCUMENTS_DIR), vectorstore)
|
||||||
|
print("Documents loaded.")
|
||||||
|
|
||||||
# Load documents – we always load; Chroma will deduplicate by ID if same content
|
# Global retriever for tool access
|
||||||
print("Loading documents into ChromaDB (if not already present)...")
|
vectorstore_retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
|
||||||
load_documents(DOCS_DIR, vectorstore)
|
|
||||||
print("Documents loaded.")
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Define tools – pass the vectorstore to the local search tool
|
# LLM and prompt
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# We wrap the tool functions to include the vectorstore argument
|
llm = ChatOllama(model="llama3")
|
||||||
|
|
||||||
def local_kb_tool(query: str, top_k: int = 3):
|
system_prompt = """You are an AI assistant that can answer questions using two sources:
|
||||||
return search_local_kb(query=query, top_k=top_k, vectorstore=vectorstore)
|
|
||||||
|
|
||||||
# Create LangChain Tool objects
|
1. A local knowledge base (ChromaDB). Use the tool ``search_local_kb`` when the
|
||||||
local_tool = Tool(
|
answer can be found in the documents.
|
||||||
name="search_local_kb",
|
2. Live web search (Tavily). Use the tool ``web_search`` when the answer requires
|
||||||
func=local_kb_tool,
|
up‑to‑date information.
|
||||||
description="Semantic search in the local knowledge base. Use when the answer is in the local documents.",
|
|
||||||
)
|
After retrieving the information, answer the user question and explicitly
|
||||||
web_tool = Tool(
|
state the source you used: either ``chromadb`` or ``tavily``.
|
||||||
name="web_search",
|
|
||||||
func=web_search,
|
If you are unsure, ask for clarification. Do not provide fabricated data.
|
||||||
description="Search the web using Tavily. Use for up‑to‑date facts or news.",
|
"""
|
||||||
)
|
|
||||||
|
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 up‑to‑date 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")
|
|
||||||
while True:
|
def answer_query(query: str) -> str:
|
||||||
try:
|
"""Return the agent's answer for *query*.
|
||||||
query = input("Query: ")
|
|
||||||
except (KeyboardInterrupt, EOFError):
|
Parameters
|
||||||
print("\nExiting.")
|
----------
|
||||||
break
|
query: str
|
||||||
if query.strip().lower() in {"exit", "quit", "q"}:
|
The user's question.
|
||||||
print("Exiting.")
|
|
||||||
break
|
Returns
|
||||||
if not query.strip():
|
-------
|
||||||
continue
|
str
|
||||||
# Run the agent
|
The agent's response.
|
||||||
try:
|
"""
|
||||||
result = agent.run(query)
|
result = agent_executor.invoke({"input": query})
|
||||||
print(f"\nAnswer:\n{result}\n")
|
return result["output"]
|
||||||
except Exception as e:
|
|
||||||
print(f"Error: {e}")
|
# ---------------------------------------------------------------------------
|
||||||
continue
|
# CLI entry point
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
if __name__ == "__main__":
|
||||||
|
print("RAG Agent ready. Type 'exit' to quit.")
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
user_input = input("\nQuery: ")
|
||||||
|
except (KeyboardInterrupt, EOFError):
|
||||||
|
print("\nExiting.")
|
||||||
|
break
|
||||||
|
if user_input.lower() in {"exit", "quit"}:
|
||||||
|
print("Goodbye!")
|
||||||
|
break
|
||||||
|
response = answer_query(user_input)
|
||||||
|
print("\nAnswer:\n", response)
|
||||||
|
|||||||
Reference in New Issue
Block a user