105 lines
3.8 KiB
Python
105 lines
3.8 KiB
Python
"""Main RAG agent implementation.
|
||
|
||
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
|
||
from pathlib import Path
|
||
|
||
from langchain_ollama import ChatOllama
|
||
from langchain.agents import AgentExecutor, create_openai_tools_agent
|
||
from langchain.prompts import ChatPromptTemplate, SystemMessagePromptTemplate, HumanMessagePromptTemplate
|
||
|
||
from vectorstore import create_vectorstore, load_documents
|
||
from rag_tools import search_local_kb, web_search
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Configuration
|
||
# ---------------------------------------------------------------------------
|
||
VECTORSTORE_DIR = Path("./chroma_db")
|
||
DOCUMENTS_DIR = Path("./documents")
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Initialise vector store and retriever
|
||
# ---------------------------------------------------------------------------
|
||
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.")
|
||
|
||
# Global retriever for tool access
|
||
vectorstore_retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# LLM and prompt
|
||
# ---------------------------------------------------------------------------
|
||
llm = ChatOllama(model="llama3")
|
||
|
||
system_prompt = """You are an AI assistant that can answer questions using two sources:
|
||
|
||
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:
|
||
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)
|