Files
task-6a1864f78a94f887e50d46da/rag_tools.py
T
2026-06-02 07:24:16 +00:00

92 lines
2.8 KiB
Python
Raw 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.
"""Tools used by the RAG agent.
Two tools are provided:
* :func:`search_local_kb` semantic search in the local ChromaDB vector store.
* :func:`web_search` realtime web search via Tavily.
Both functions are decorated with :func:`langchain.tools.tool` so that LangChain
can expose them to the agent.
"""
from __future__ import annotations
import os
from typing import List
from langchain.chains import RetrievalQA
from langchain.chroma import Chroma
from langchain_ollama import ChatOllama
from langchain_tavily import TavilySearchResults
from langchain.tools import tool
# Load environment variables (TAVILY_API_KEY)
from dotenv import load_dotenv
load_dotenv()
# ---------------------------------------------------------------------------
# Local KB search
# ---------------------------------------------------------------------------
@tool("search_local_kb")
def search_local_kb(query: str, top_k: int = 3, vectorstore: Chroma | None = None) -> str:
"""Return the best answer from the local ChromaDB store.
Parameters
----------
query: str
The user question.
top_k: int
Number of documents to retrieve.
vectorstore: Chroma | None
If ``None`` the function will attempt to load the default store.
Returns
-------
str
The answer prefixed with the source identifier.
"""
if vectorstore is None:
# Load the default store
from vectorstore import create_vectorstore
vectorstore = create_vectorstore()
retriever = vectorstore.as_retriever(search_kwargs={"k": top_k})
qa_chain = RetrievalQA.from_chain_type(
llm=ChatOllama(model="llama3", temperature=0),
chain_type="stuff",
retriever=retriever,
)
result = qa_chain.run(query)
return f"[chromadb] {result}"
# ---------------------------------------------------------------------------
# Web search via Tavily
# ---------------------------------------------------------------------------
@tool("web_search")
def web_search(query: str, top_k: int = 3) -> str:
"""Return an answer based on a Tavily web search.
Parameters
----------
query: str
The user question.
top_k: int
Number of search results to retrieve.
Returns
-------
str
The answer prefixed with the source identifier.
"""
api_key = os.getenv("TAVILY_API_KEY")
if not api_key:
raise RuntimeError("TAVILY_API_KEY environment variable is not set")
tavily = TavilySearchResults(api_key=api_key, top_k=top_k)
results = tavily.run(query)
snippets = "\n\n".join(item.get("content", "") for item in results)
llm = ChatOllama(model="llama3", temperature=0)
prompt = f"Answer the question based on the following web snippets:\n{snippets}\n\nQuestion: {query}\nAnswer:" # noqa: E501
answer = llm.invoke(prompt).content
return f"[tavily] {answer}"