Files
2026-05-28 16:57:45 +00:00

73 lines
2.7 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
Simple CLI for the RAGagent.
The agent automatically chooses between a local ChromaDB search and a web search via Tavily.
It prints the answer together with the source label.
"""
import os
from pathlib import Path
from dotenv import load_dotenv
from langchain_ollama import OllamaLLM
from langchain_tavily import TavilySearchResults
from langchain.tools import tool
from langchain.agents import initialize_agent, AgentType
from langchain.chains import RetrievalQA
from vectorstore import create_vectorstore, load_documents
# Load env for Tavily API key
load_dotenv()
TAVILY_API_KEY = os.getenv("TAVILY_API_KEY")
if not TAVILY_API_KEY:
raise RuntimeError("TAVILY_API_KEY is missing in .env")
# 1. Vector store and retriever
vectorstore = create_vectorstore()
load_documents("documents", vectorstore)
retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
# 2. Tools
@tool(name="search_local_kb", description="Search the local knowledge base using ChromaDB.")
def search_local_kb(query: str, top_k: int = 3) -> str:
docs = retriever.invoke({"query": query}) if hasattr(retriever, "invoke") else retriever.get_relevant_documents(query)
return "\n---\n".join([f"{d.metadata.get('source')}:\n{d.page_content[:200]}" for d in docs])
@tool(name="web_search", description="Search the web using Tavily.")
def web_search(query: str) -> str:
tavily = TavilySearchResults(api_key=TAVILY_API_KEY, max_results=3)
results = tavily.run(query)
return "\n---\n".join([f"{r['title']} ({r['url']}):\n{r.get('content', '')[:200]}" for r in results])
# 3. Agent with simple routing logic
from langchain.prompts import ChatPromptTemplate, SystemMessagePromptTemplate, HumanMessagePromptTemplate
from langchain.chat_models import ChatOllama
chat = ChatOllama(model="llama3")
system_prompt = "You are an assistant that can answer questions using either a local knowledge base or the web. If the question is about recent events or news, use web_search; otherwise use search_local_kb. Return the answer and specify the source as either chromadb or tavily."
prompt = ChatPromptTemplate.from_messages([
SystemMessagePromptTemplate.from_template(system_prompt),
HumanMessagePromptTemplate.from_template("{input}")
])
agent_chain = initialize_agent(
tools=[search_local_kb, web_search],
llm=chat,
agent=AgentType.CHAT_ZERO_SHOT_REACT_DESCRIPTION,
verbose=True,
)
# 4. CLI loop
if __name__ == "__main__":
print("RAG Agent ready. Type 'exit' to quit.")
while True:
try:
q = input("Query: ")
except EOFError:
break
if q.strip().lower() in {"exit", "quit"}:
break
response = agent_chain.run(q)
print(f"\nAnswer:\n{response}\n")