From b9296a06f30492ad01096f6bb98a94570c2dd608 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=9A=D0=B8=D1=80=D0=B8=D0=BB=D0=BB=20=D0=A0=D0=BE=D0=BC?= =?UTF-8?q?=D0=B0=D0=BD=D0=BE=D0=B2?= Date: Tue, 2 Jun 2026 07:10:50 +0000 Subject: [PATCH] Update agent.py --- agent.py | 155 ++++++++++++++++++++++++++++++++----------------------- 1 file changed, 90 insertions(+), 65 deletions(-) diff --git a/agent.py b/agent.py index 1aa1828..a7997dd 100644 --- a/agent.py +++ b/agent.py @@ -1,74 +1,99 @@ -""" -Agent setup with tools for local ChromaDB search and Tavily web search. +"""Main script for the RAG agent with ChromaDB and Tavily. + +The script: +1. Loads or creates the Chroma vector store. +2. Loads documents from the `documents/` folder. +3. Sets up the LangChain agent with two tools: `search_local_kb` and `web_search`. +4. Runs a simple CLI loop. """ import os -from typing import List +from pathlib import Path from langchain_ollama import ChatOllama -from langchain.tools import tool -from langchain.schema import Document -from langchain_chroma import Chroma -from langchain_ollama import OllamaEmbeddings -from langchain_text_splitters import RecursiveCharacterTextSplitter -from langchain_community.document_loaders import TextLoader -from langchain_community.document_loaders import MarkdownLoader - -from tavily import TavilySearchResults - -# Load environment variables -from dotenv import load_dotenv -load_dotenv() - -# Load vectorstore -from vectorstore import create_vectorstore, load_documents - -# Persist directory -PERSIST_DIR = "./chroma_db" - -# Create or load vectorstore -vectorstore = create_vectorstore(persist_directory=PERSIST_DIR) - -# Load documents from documents folder if not already loaded -if not os.path.exists(PERSIST_DIR) or not os.listdir(PERSIST_DIR): - print("Loading documents into vector store...") - load_documents("documents", vectorstore) - -# Define tools -@tool -def search_local_kb(query: str, top_k: int = 3) -> str: - """Semantic search in local ChromaDB knowledge base.""" - retriever = vectorstore.as_retriever(search_kwargs={"k": top_k}) - docs = retriever.get_relevant_documents(query) - if not docs: - return "No relevant documents found in local knowledge base." - # Concatenate content - content = "\n\n".join([f"Source: {doc.metadata.get('source', 'unknown')}\n{doc.page_content}" for doc in docs]) - return f"[Local KB]\n{content}" - -@tool -def web_search(query: str) -> str: - """Web search using Tavily.""" - tavily = TavilySearchResults(api_key=os.getenv("TAVILY_API_KEY")) - results = tavily.run(query) - if not results: - return "No web results found." - # Format results - formatted = "\n\n".join([f"{i+1}. {r.get('title', 'No title')}\n{r.get('url', '')}\n{r.get('content', '')}" for i, r in enumerate(results)]) - return f"[Web Search]\n{formatted}" - -# Create agent -llm = ChatOllama(model="llama3") - from langchain.agents import initialize_agent, AgentType +from langchain.tools import Tool -agent_executor = initialize_agent( - tools=[search_local_kb, web_search], - llm=llm, - agent=AgentType.ZERO_SHOT_REACT_DESCRIPTION, - verbose=True, - handle_parsing_errors=True, +from vectorstore import create_vectorstore, load_documents +from rag_tools import search_local_kb, web_search + +# --------------------------------------------------------------------------- +# Configuration +# --------------------------------------------------------------------------- +CHROMA_DIR = "./chroma_db" +DOCS_DIR = "./documents" +MODEL = "llama3" + +# --------------------------------------------------------------------------- +# Helper: load or create vector store +# --------------------------------------------------------------------------- +vectorstore = create_vectorstore(persist_directory=CHROMA_DIR) + +# Load documents – we always load; Chroma will deduplicate by ID if same content +print("Loading documents into ChromaDB (if not already present)...") +load_documents(DOCS_DIR, vectorstore) +print("Documents loaded.") + +# --------------------------------------------------------------------------- +# Define tools – pass the vectorstore to the local search tool +# --------------------------------------------------------------------------- +# 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) + +# Create LangChain Tool objects +local_tool = Tool( + name="search_local_kb", + func=local_kb_tool, + description="Semantic search in the local knowledge base. Use when the answer is in the local documents.", +) +web_tool = Tool( + name="web_search", + func=web_search, + description="Search the web using Tavily. Use for up‑to‑date facts or news.", ) -# Expose agent_executor -__all__ = ["agent_executor"] \ No newline at end of file +# --------------------------------------------------------------------------- +# Agent setup +# --------------------------------------------------------------------------- +llm = ChatOllama(model=MODEL, temperature=0.0) + +system_prompt = ( + "You are an assistant that answers user questions. " + "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 +# --------------------------------------------------------------------------- +print("\nRAG Agent ready. Type your question (or 'exit' to quit).\n") +while True: + try: + query = input("Query: ") + except (KeyboardInterrupt, EOFError): + print("\nExiting.") + break + if query.strip().lower() in {"exit", "quit", "q"}: + print("Exiting.") + break + if not query.strip(): + continue + # Run the agent + try: + result = agent.run(query) + print(f"\nAnswer:\n{result}\n") + except Exception as e: + print(f"Error: {e}") + continue