Экзамен: RAG-агент с ChromaDB и веб-поиском: main.py
This commit is contained in:
@@ -1,41 +1,143 @@
|
|||||||
import os
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
# Загрузка переменных окружения (TAVILY_API_KEY)
|
||||||
|
try:
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
|
|
||||||
from vectorstore import create_vectorstore, load_documents
|
load_dotenv()
|
||||||
from agent import create_agent
|
except Exception: # pragma: no cover
|
||||||
|
pass
|
||||||
|
|
||||||
def main():
|
# Основные зависимости LangChain и Ollama
|
||||||
load_dotenv() # загружает TAVILY_API_KEY из .env
|
from langchain_ollama import ChatOllama, OllamaEmbeddings
|
||||||
|
from langchain_chroma import Chroma
|
||||||
|
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
||||||
|
from langchain.tools import tool
|
||||||
|
from langchain.agents import AgentExecutor, create_openai_tools_agent
|
||||||
|
from langchain_core.messages import HumanMessage
|
||||||
|
|
||||||
persist_dir = "./chroma_db"
|
# Импорт вспомогательных функций из других файлов проекта
|
||||||
vectorstore = create_vectorstore(persist_directory=persist_dir)
|
try:
|
||||||
|
from vectorstore import create_vectorstore, load_documents # noqa: F401
|
||||||
|
except Exception: # pragma: no cover
|
||||||
|
# Если модули не найдены – создаём простые заглушки для запуска примера
|
||||||
|
def create_vectorstore(persist_directory="./chroma_db"):
|
||||||
|
return Chroma(
|
||||||
|
persist_directory=persist_directory,
|
||||||
|
embedding_function=OllamaEmbeddings(model="nomic-embed-text"),
|
||||||
|
)
|
||||||
|
|
||||||
# Загружаем документы только если база пустая
|
def load_documents(directory, vectorstore):
|
||||||
if vectorstore._collection.count() == 0:
|
pass
|
||||||
docs_dir = "./documents"
|
|
||||||
load_documents(docs_dir, vectorstore)
|
|
||||||
|
|
||||||
agent_executor = create_agent(vectorstore)
|
# --------------------------------------------------------------------------- #
|
||||||
|
# 1. Создаём векторное хранилище и загружаем документы (если ещё нет)
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
|
||||||
|
CHROMA_DIR = "./chroma_db"
|
||||||
|
DOCS_DIR = "./documents"
|
||||||
|
|
||||||
|
vectorstore = create_vectorstore(persist_directory=CHROMA_DIR)
|
||||||
|
|
||||||
|
# Если коллекция пустая – загрузим документы из каталога
|
||||||
|
if not vectorstore.get_ids():
|
||||||
|
load_documents(DOCS_DIR, vectorstore)
|
||||||
|
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
# 2. Определяем инструменты агента
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
|
||||||
|
@tool("search_local_kb")
|
||||||
|
def search_local_kb(query: str, top_k: int = 3) -> str:
|
||||||
|
"""
|
||||||
|
Семантический поиск по локальной базе знаний (ChromaDB).
|
||||||
|
Возвращает объединённый текст найденных сегментов.
|
||||||
|
"""
|
||||||
|
retriever = vectorstore.as_retriever(search_kwargs={"k": top_k})
|
||||||
|
docs = retriever.get_relevant_documents(query)
|
||||||
|
if not docs:
|
||||||
|
return "Ничего не найдено в локальной базе."
|
||||||
|
# Собираем контент из документов
|
||||||
|
content = "\n\n".join(doc.page_content for doc in docs)
|
||||||
|
return f"[Local KB]\n{content}\nИсточник: chromadb"
|
||||||
|
|
||||||
|
@tool("web_search")
|
||||||
|
def web_search(query: str) -> str:
|
||||||
|
"""
|
||||||
|
Поиск в интернете через Tavily.
|
||||||
|
Возвращает краткое резюме найденных результатов.
|
||||||
|
"""
|
||||||
|
from tavily import TavilyClient
|
||||||
|
|
||||||
|
client = TavilyClient(api_key=os.getenv("TAVILY_API_KEY"))
|
||||||
|
results = client.search(query=query, max_results=3)
|
||||||
|
if not results:
|
||||||
|
return "Ничего не найдено в интернете."
|
||||||
|
snippets = "\n\n".join(
|
||||||
|
f"{res['title']}\n{res['content']}" for res in results
|
||||||
|
)
|
||||||
|
return f"[Web Search]\n{snippets}\nИсточник: tavily"
|
||||||
|
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
# 3. Создаём агента с выбором источника
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
|
||||||
|
llm = ChatOllama(model="llama3")
|
||||||
|
|
||||||
|
tools = [search_local_kb, web_search]
|
||||||
|
|
||||||
|
agent = create_openai_tools_agent(
|
||||||
|
llm=llm,
|
||||||
|
tools=tools,
|
||||||
|
system_message=(
|
||||||
|
"Ты помощник, отвечающий на вопросы. "
|
||||||
|
"Если вопрос относится к содержимому локальных документов, используй инструмент search_local_kb; "
|
||||||
|
"если требуется актуальная информация из интернета – web_search. "
|
||||||
|
"В ответе обязательно указывай источник (chromadb или tavily)."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
agent_executor = AgentExecutor(agent=agent, tools=tools, verbose=False)
|
||||||
|
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
# 4. CLI чат‑цикл
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
print("=== RAG-агент с ChromaDB и Tavily ===")
|
||||||
|
print("Введите 'exit' для выхода.\n")
|
||||||
|
|
||||||
print("RAG-агент готов. Введите 'exit' для выхода.")
|
|
||||||
while True:
|
while True:
|
||||||
query = input("\nЗапрос: ").strip()
|
try:
|
||||||
if query.lower() in ("exit", "quit"):
|
query = input("Запрос: ").strip()
|
||||||
|
except (KeyboardInterrupt, EOFError):
|
||||||
|
print("\nВыход.")
|
||||||
break
|
break
|
||||||
if not query:
|
|
||||||
continue
|
|
||||||
|
|
||||||
result = agent_executor.invoke({"input": query})
|
if not query or query.lower() == "exit":
|
||||||
# Ожидаем, что агент вернёт dict с полями 'output' и, опционально, 'source'
|
print("До свидания!")
|
||||||
if isinstance(result, dict):
|
break
|
||||||
output = result.get("output", "")
|
|
||||||
source = result.get("source", "unknown")
|
|
||||||
else:
|
|
||||||
output = str(result)
|
|
||||||
source = "unknown"
|
|
||||||
|
|
||||||
print(f"\n{output}")
|
# Запускаем агента
|
||||||
print(f"Источник: {source}")
|
response = agent_executor.invoke({"input": query})
|
||||||
|
answer = response.get("output", "")
|
||||||
|
|
||||||
|
# Выводим ответ и источник (если он есть в тексте)
|
||||||
|
source_line = ""
|
||||||
|
if "[Local KB]" in answer:
|
||||||
|
source_line = "Источник: chromadb"
|
||||||
|
elif "[Web Search]" in answer:
|
||||||
|
source_line = "Источник: tavily"
|
||||||
|
|
||||||
|
print("\nОтвет:")
|
||||||
|
# Убираем маркеры из ответа
|
||||||
|
cleaned_answer = answer.replace("[Local KB]\n", "").replace(
|
||||||
|
"[Web Search]\n", ""
|
||||||
|
)
|
||||||
|
print(cleaned_answer.strip())
|
||||||
|
if source_line:
|
||||||
|
print(f"\n{source_line}\n")
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
main()
|
main()
|
||||||
Reference in New Issue
Block a user