Add agent.py
This commit is contained in:
@@ -0,0 +1,74 @@
|
|||||||
|
"""
|
||||||
|
Agent setup with tools for local ChromaDB search and Tavily web search.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
from typing import List
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
agent_executor = initialize_agent(
|
||||||
|
tools=[search_local_kb, web_search],
|
||||||
|
llm=llm,
|
||||||
|
agent=AgentType.ZERO_SHOT_REACT_DESCRIPTION,
|
||||||
|
verbose=True,
|
||||||
|
handle_parsing_errors=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Expose agent_executor
|
||||||
|
__all__ = ["agent_executor"]
|
||||||
Reference in New Issue
Block a user