Files
task-6a1864f78a94f887e50d46da/agent.py
T
2026-06-01 17:40:29 +00:00

74 lines
2.3 KiB
Python

"""
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"]