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

81 lines
2.5 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 for the RAG agent.
This module defines two LangChain tools:
* ``search_local_kb`` semantic search in the local ChromaDB vector store.
* ``web_search`` realtime web search using Tavily.
Both tools return a string containing the retrieved information.
"""
from typing import List
from langchain_ollama import ChatOllama
from langchain_tavily import TavilySearchResults
from langchain.tools import tool
# The LLM used for summarising or formatting responses
llm = ChatOllama(model="llama3")
# Tavily client the API key is read from the environment by the package
# (requires a .env file or the TAVILY_API_KEY environment variable).
search = TavilySearchResults()
# ---------------------------------------------------------------------------
# Local knowledge base search tool
# ---------------------------------------------------------------------------
@tool("search_local_kb")
def search_local_kb(query: str, top_k: int = 3) -> str:
"""Perform a semantic search in the local ChromaDB vector store.
Parameters
----------
query: str
The user question.
top_k: int, optional
Number of top documents to return. Defaults to 3.
Returns
-------
str
A formatted string containing the retrieved passages.
"""
# The vectorstore is expected to be loaded globally the agent will
# provide it via the tool context. We simply call the retriever.
retriever = globals().get("vectorstore_retriever")
if retriever is None:
raise RuntimeError("Vector store retriever not configured for the tool.")
docs = retriever.get_relevant_documents(query, k=top_k)
# Concatenate the documents into a single string.
passages = "\n\n".join(doc.page_content for doc in docs)
return passages
# ---------------------------------------------------------------------------
# Web search tool
# ---------------------------------------------------------------------------
@tool("web_search")
def web_search(query: str) -> str:
"""Search the web using Tavily and return the top results.
Parameters
----------
query: str
The user question.
Returns
-------
str
A formatted string containing the search results.
"""
results = search.run(query)
# TavilySearchResults returns a list of dicts with keys: title, url, content
formatted = []
for r in results:
formatted.append(f"Title: {r.get('title', 'N/A')}\nURL: {r.get('url', 'N/A')}\nSnippet: {r.get('content', 'N/A')}\n")
return "\n\n".join(formatted)
# End of rag_tools.py