add main.py
This commit is contained in:
@@ -0,0 +1,157 @@
|
||||
"""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())
|
||||
Reference in New Issue
Block a user