Files
Agent/local_knowledge_base.py

156 lines
6.0 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.
import os
from pathlib import Path
from qdrant_client import QdrantClient
from qdrant_client.models import Distance, VectorParams, PointStruct
from langchain_community.document_loaders import (
TextLoader,
PyPDFLoader,
DirectoryLoader,
)
from langchain_text_splitters import Language, RecursiveCharacterTextSplitter
from langchain_community.embeddings import OllamaEmbeddings
def should_skip(path: str) -> bool:
"""Проверяет, нужно ли пропустить путь (системные папки)"""
skip_dirs = {'.git', 'build', 'vendor', 'node_modules', '__pycache__', '.venv', 'venv'}
return any(part in skip_dirs for part in Path(path).parts)
class LocalKnowledgeBase:
def __init__(self, collection_name="documents", persist_dir="./qdrant_data"):
# Инициализируем клиент Qdrant
self.client = QdrantClient(
path=persist_dir, prefer_grpc=True # Локальное хранение
)
self.collection_name = collection_name
# Инициализация эмбеддингов
self.embeddings = OllamaEmbeddings(
model="nomic-embed-text", # Хорошие локальные эмбеддинги
)
# Создаем коллекцию если её нет
self._create_collection()
def _create_collection(self):
try:
self.client.get_collection(self.collection_name)
print(f"Коллекция {self.collection_name} уже существует")
except:
# Создаем новую коллекцию
self.client.create_collection(
collection_name=self.collection_name,
vectors_config=VectorParams(
size=768, # Размерность эмбеддингов nomic-embed-text
distance=Distance.COSINE,
),
)
print(f"Создана коллекция {self.collection_name}")
def load_documents(self, directory_path: str):
"""Загрузка документов из директории (рекурсивно)"""
loaders = {
".pdf": PyPDFLoader,
".txt": lambda path: TextLoader(path, encoding="utf-8"),
".docx": lambda path: DirectoryLoader(path, glob="**/*.docx"),
".hpp": lambda path: TextLoader(path, encoding="utf-8"),
".cpp": lambda path: TextLoader(path, encoding="utf-8"),
".h": lambda path: TextLoader(path, encoding="utf-8"),
".cc": lambda path: TextLoader(path, encoding="utf-8"),
}
all_documents = []
# Рекурсивный обход с фильтрацией папок
for root, dirs, files in os.walk(directory_path):
# Фильтруем директории на месте
dirs[:] = [d for d in dirs if not should_skip(os.path.join(root, d))]
for file_path in files:
ext = os.path.splitext(file_path)[1]
if ext in loaders:
full_path = os.path.join(root, file_path)
try:
loader = loaders[ext](full_path)
documents = loader.load()
all_documents.extend(documents)
print(f"Загружен {file_path}: {len(documents)} документов")
except Exception as e:
print(f"Ошибка загрузки {file_path}: {e}")
# Разбиваем на чанки
text_splitter = RecursiveCharacterTextSplitter.from_language(
language=Language.CPP,
chunk_size=1000,
chunk_overlap=200,
length_function=len,
)
chunks = text_splitter.split_documents(all_documents)
print(f"Всего чанков: {len(chunks)}")
# Создаем эмбеддинги и сохраняем в Qdrant
self._index_documents(chunks)
return chunks
def _index_documents(self, documents):
"""Индексация документов в Qdrant"""
points = []
for i, doc in enumerate(documents):
# Создаем эмбеддинг для каждого чанка
embedding = self.embeddings.embed_query(doc.page_content)
point = PointStruct(
id=i,
vector=embedding,
payload={
"text": doc.page_content,
"source": doc.metadata.get("source", "unknown"),
"page": doc.metadata.get("page", 0),
"file_type": "cpp",
},
)
points.append(point)
# Пакетная загрузка каждые 100 точек
if len(points) >= 100:
self.client.upsert(collection_name=self.collection_name, points=points)
points = []
print(f"Индексировано {i+1} документов")
# Загружаем оставшиеся
if points:
self.client.upsert(collection_name=self.collection_name, points=points)
print(f"Индексация завершена. Всего документов: {len(documents)}")
def search(self, query: str, top_k: int = 5):
"""Поиск в базе знаний"""
# Создаем эмбеддинг запроса
query_embedding = self.embeddings.embed_query(query)
# Ищем в Qdrant
search_result = self.client.query_points(
collection_name=self.collection_name,
query=query_embedding,
limit=top_k,
)
# Форматируем результаты
results = []
for hit in search_result.points:
results.append(
{
"text": hit.payload["text"],
"score": hit.score,
"source": hit.payload.get("source", "unknown"),
}
)
return results