158 lines
5.2 KiB
Python
158 lines
5.2 KiB
Python
"""RAG-агент: ChromaDB (локально) + Tavily (веб), маршрутизация через create_agent."""
|
|
from __future__ import annotations
|
|
|
|
import os
|
|
import sys
|
|
from pathlib import Path
|
|
|
|
from dotenv import load_dotenv
|
|
from langchain.agents import create_agent
|
|
from langchain_core.tools import tool
|
|
from langchain_ollama import ChatOllama
|
|
from tavily import TavilyClient
|
|
|
|
from vectorstore import create_vectorstore, load_documents
|
|
|
|
load_dotenv()
|
|
|
|
DEFAULT_OLLAMA_MODEL = "llama3"
|
|
DOCUMENTS_DIR = Path(__file__).resolve().parent / "documents"
|
|
CHROMA_DIR = Path(__file__).resolve().parent / "chroma_db"
|
|
|
|
_vectorstore = None
|
|
|
|
|
|
def get_vectorstore():
|
|
global _vectorstore
|
|
if _vectorstore is None:
|
|
_vectorstore = create_vectorstore(persist_directory=str(CHROMA_DIR))
|
|
return _vectorstore
|
|
|
|
|
|
def init_knowledge_base() -> int:
|
|
vs = get_vectorstore()
|
|
count = load_documents(str(DOCUMENTS_DIR), vs)
|
|
print(f"Загружено чанков в ChromaDB: {count}")
|
|
return count
|
|
|
|
|
|
@tool
|
|
def search_local_kb(query: str, top_k: int = 4) -> str:
|
|
"""Семантический поиск по локальной базе знаний (ChromaDB).
|
|
|
|
Используй для вопросов о конспектах, LangGraph, материалах курса.
|
|
"""
|
|
retriever = get_vectorstore().as_retriever(search_kwargs={"k": top_k})
|
|
docs = retriever.invoke(query)
|
|
if not docs:
|
|
return "[Local KB] Ничего не найдено в chromadb."
|
|
parts = []
|
|
for i, doc in enumerate(docs, 1):
|
|
src = doc.metadata.get("source", "unknown")
|
|
parts.append(f"[{i}] ({src}) {doc.page_content[:500]}")
|
|
return "[Local KB]\n" + "\n\n".join(parts)
|
|
|
|
|
|
@tool
|
|
def web_search(query: str) -> str:
|
|
"""Поиск актуальной информации в интернете через Tavily.
|
|
|
|
Используй для новостей, свежих фактов и событий вне локальных конспектов.
|
|
"""
|
|
api_key = os.getenv("TAVILY_API_KEY", "")
|
|
if not api_key:
|
|
return "[Web Search] TAVILY_API_KEY не задан в .env"
|
|
client = TavilyClient(api_key=api_key)
|
|
resp = client.search(query, max_results=5)
|
|
results = resp.get("results", [])
|
|
if not results:
|
|
return "[Web Search] Ничего не найдено (tavily)."
|
|
lines = []
|
|
for r in results:
|
|
lines.append(
|
|
f"- {r.get('title', '')}\n url: {r.get('url', '')}\n {r.get('content', '')[:300]}"
|
|
)
|
|
return "[Web Search]\n" + "\n".join(lines)
|
|
|
|
|
|
SYSTEM_PROMPT = """Ты RAG-агент с двумя инструментами:
|
|
- search_local_kb — локальные конспекты (ChromaDB)
|
|
- web_search — актуальная информация из интернета (Tavily)
|
|
|
|
Правила:
|
|
1. Вопросы про LangGraph, конспекты, материалы курса, «наши документы» → search_local_kb.
|
|
2. Вопросы про последние новости, актуальные события, свежие факты из сети → web_search.
|
|
3. Не вызывай оба инструмента без необходимости.
|
|
4. В конце ответа укажи строку: Источник: chromadb | tavily
|
|
"""
|
|
|
|
|
|
def build_llm() -> ChatOllama:
|
|
return ChatOllama(model=os.getenv("OLLAMA_MODEL", DEFAULT_OLLAMA_MODEL), temperature=0.2)
|
|
|
|
|
|
def build_agent():
|
|
llm = build_llm()
|
|
tools = [search_local_kb, web_search]
|
|
return create_agent(model=llm, tools=tools, system_prompt=SYSTEM_PROMPT)
|
|
|
|
|
|
def _extract_answer(result) -> str:
|
|
messages = result.get("messages", [])
|
|
if not messages:
|
|
return str(result)
|
|
last = messages[-1]
|
|
return getattr(last, "content", str(last))
|
|
|
|
|
|
def chat_loop() -> None:
|
|
agent = build_agent()
|
|
print("RAG-агент (ChromaDB + Tavily). Команды: exit / quit")
|
|
print("Перед первым запуском: python main.py --init\n")
|
|
while True:
|
|
try:
|
|
user = input("Вы: ").strip()
|
|
except (EOFError, KeyboardInterrupt):
|
|
print("\nВыход.")
|
|
break
|
|
if not user:
|
|
continue
|
|
if user.lower() in {"exit", "quit", "q"}:
|
|
print("Выход.")
|
|
break
|
|
result = agent.invoke({"messages": [{"role": "user", "content": user}]})
|
|
answer = _extract_answer(result)
|
|
print(f"\nАгент:\n{answer}\n")
|
|
|
|
|
|
def demo() -> None:
|
|
agent = build_agent()
|
|
samples = [
|
|
"Что в наших конспектах про LangGraph?",
|
|
"Какие последние новости про AI-агентов?",
|
|
]
|
|
for q in samples:
|
|
print(f"Запрос: {q}")
|
|
result = agent.invoke({"messages": [{"role": "user", "content": q}]})
|
|
print(_extract_answer(result))
|
|
print("-" * 60)
|
|
|
|
|
|
def main() -> int:
|
|
if "--init" in sys.argv:
|
|
init_knowledge_base()
|
|
return 0
|
|
if "--demo" in sys.argv:
|
|
if not CHROMA_DIR.exists() or not any(CHROMA_DIR.iterdir()):
|
|
init_knowledge_base()
|
|
demo()
|
|
return 0
|
|
if not CHROMA_DIR.exists() or not any(CHROMA_DIR.iterdir()):
|
|
init_knowledge_base()
|
|
chat_loop()
|
|
return 0
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|