diff --git a/agent.py b/agent.py new file mode 100644 index 0000000..4a52e12 --- /dev/null +++ b/agent.py @@ -0,0 +1,68 @@ +import os +from langchain_ollama import ChatOllama +from langchain.agents import initialize_agent, AgentType +from langchain.tools import tool +from langchain_tavily import TavilySearchResults +from vectorstore import VectorStore + +# Global vector store instance to be set by get_agent +vector_store: VectorStore | None = None + +@tool +def local_kb_search(query: str) -> str: + """Search local knowledge base (ChromaDB) for relevant information.""" + if vector_store is None: + return "Vector store not initialized." + docs = vector_store.get_store().similarity_search(query, k=3) + if not docs: + return "No relevant documents found." + answer = "\n".join([doc.page_content.strip() for doc in docs]) + return f"{answer}\nSource: chromadb" + +@tool +def web_search(query: str) -> str: + """Search the web using Tavily.""" + api_key = os.getenv("TAVILY_API_KEY") + if not api_key: + return "TAVILY_API_KEY not set. Please set the environment variable." + try: + tavily = TavilySearchResults(tavily_api_key=api_key) + results = tavily.run({"query": query}) + if not results: + return "No results found." + answer = "" + for i, res in enumerate(results[:3]): + title = res.get("title") or res.get("name") or "No title" + url = res.get("url") or "" + content = res.get("content") or "" + answer += f"{i+1}. {title} ({url})\n{content}\n\n" + return f"{answer}\nSource: tavily" + except Exception as e: + return f"Error during web search: {e}" + + +def get_agent(vectorstore: VectorStore): + global vector_store + vector_store = vectorstore + base_url = os.getenv("CHAT_BASE_URL") + api_key = os.getenv("CHAT_API_KEY") + model = os.getenv("CHAT_MODEL", "llama3") + if not base_url: + raise EnvironmentError("CHAT_BASE_URL not set. Please set the environment variable.") + if not api_key: + raise EnvironmentError("CHAT_API_KEY not set. Please set the environment variable.") + llm = ChatOllama( + temperature=0, + model=model, + base_url=base_url, + api_key=api_key + ) + tools = [local_kb_search, web_search] + agent = initialize_agent( + tools, + llm, + agent=AgentType.OPENAI_FUNCTIONS, + verbose=True, + handle_parsing_errors=True + ) + return agent