Update agent.py
This commit is contained in:
@@ -1,104 +1,57 @@
|
|||||||
"""Main RAG agent implementation.
|
"""
|
||||||
|
Main agent logic: decides whether to use local KB or web search.
|
||||||
The agent can answer questions using either the local ChromaDB knowledge base
|
|
||||||
or live web search via Tavily. The decision of which tool to use is made by
|
|
||||||
the LLM itself based on the prompt.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import os
|
import os
|
||||||
from pathlib import Path
|
from typing import Dict, Any
|
||||||
|
|
||||||
from langchain_ollama import ChatOllama
|
from langchain_ollama import ChatOllama
|
||||||
from langchain.agents import AgentExecutor, create_openai_tools_agent
|
from langchain.agents import initialize_agent, AgentType, Tool, AgentExecutor
|
||||||
from langchain.prompts import ChatPromptTemplate, SystemMessagePromptTemplate, HumanMessagePromptTemplate
|
from langchain_core.messages import HumanMessage
|
||||||
|
|
||||||
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
|
||||||
|
from vectorstore import create_vectorstore, load_documents
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# Load or create vector store
|
||||||
# Configuration
|
vectorstore = create_vectorstore()
|
||||||
# ---------------------------------------------------------------------------
|
# Load documents from the documents folder if not already loaded
|
||||||
VECTORSTORE_DIR = Path("./chroma_db")
|
if not vectorstore._collection.count(): # type: ignore[attr-defined]
|
||||||
DOCUMENTS_DIR = Path("./documents")
|
load_documents("./documents", vectorstore)
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# Define tools
|
||||||
# Initialise vector store and retriever
|
tools = [
|
||||||
# ---------------------------------------------------------------------------
|
Tool(name="search_local_kb", func=search_local_kb, description="Search the local knowledge base."),
|
||||||
vectorstore = create_vectorstore(str(VECTORSTORE_DIR))
|
Tool(name="web_search", func=web_search, description="Search the web using Tavily."),
|
||||||
# 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.")
|
|
||||||
|
|
||||||
# Global retriever for tool access
|
# System prompt to guide the agent
|
||||||
vectorstore_retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
|
system_prompt = (
|
||||||
|
"You are an AI assistant. For questions about local documents use the 'search_local_kb' tool. "
|
||||||
|
"For recent news or facts not in the local docs, use 'web_search'. "
|
||||||
|
"Always indicate the source of your answer (chromadb or tavily)."
|
||||||
|
)
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# Create the agent executor
|
||||||
# LLM and prompt
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
llm = ChatOllama(model="llama3")
|
llm = ChatOllama(model="llama3")
|
||||||
|
agent_executor = initialize_agent(
|
||||||
|
tools=tools,
|
||||||
|
llm=llm,
|
||||||
|
agent=AgentType.OPENAI_FUNCTIONS,
|
||||||
|
verbose=True,
|
||||||
|
system_message=system_prompt,
|
||||||
|
)
|
||||||
|
|
||||||
system_prompt = """You are an AI assistant that can answer questions using two sources:
|
def main():
|
||||||
|
print("Welcome to the RAG agent. Type 'exit' to quit.")
|
||||||
1. A local knowledge base (ChromaDB). Use the tool ``search_local_kb`` when the
|
|
||||||
answer can be found in the documents.
|
|
||||||
2. Live web search (Tavily). Use the tool ``web_search`` when the answer requires
|
|
||||||
up‑to‑date information.
|
|
||||||
|
|
||||||
After retrieving the information, answer the user question and explicitly
|
|
||||||
state the source you used: either ``chromadb`` or ``tavily``.
|
|
||||||
|
|
||||||
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
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Tools are automatically discovered via the @tool decorator in rag_tools.py
|
|
||||||
tools = [search_local_kb, web_search]
|
|
||||||
|
|
||||||
agent = create_openai_tools_agent(llm=llm, tools=tools, prompt=prompt)
|
|
||||||
agent_executor = AgentExecutor.from_agent_and_tools(agent=agent, tools=tools, verbose=True)
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Public API
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
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:
|
user_input = input("\nUser: ")
|
||||||
user_input = input("\nQuery: ")
|
|
||||||
except (KeyboardInterrupt, EOFError):
|
|
||||||
print("\nExiting.")
|
|
||||||
break
|
|
||||||
if user_input.lower() in {"exit", "quit"}:
|
if user_input.lower() in {"exit", "quit"}:
|
||||||
print("Goodbye!")
|
print("Goodbye!")
|
||||||
break
|
break
|
||||||
response = answer_query(user_input)
|
# Run the agent
|
||||||
print("\nAnswer:\n", response)
|
result = agent_executor.invoke({"input": user_input})
|
||||||
|
# The result may contain tool calls and final answer
|
||||||
|
print("\nAssistant:", result.get("output", ""))
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
|
|||||||
Reference in New Issue
Block a user