204 lines
8.6 KiB
Python
204 lines
8.6 KiB
Python
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, project_name: str = "default", persist_dir="./qdrant_data"):
|
||
# Инициализируем клиент Qdrant
|
||
self.client = QdrantClient(
|
||
path=persist_dir, prefer_grpc=True # Локальное хранение
|
||
)
|
||
|
||
self.persist_dir = persist_dir
|
||
self.set_project(project_name)
|
||
|
||
# Инициализация эмбеддингов
|
||
self.embeddings = OllamaEmbeddings(
|
||
model="nomic-embed-text", # Хорошие локальные эмбеддинги
|
||
)
|
||
|
||
def set_project(self, project_name: str):
|
||
"""Переключиться на проект (создать/выбрать коллекцию)"""
|
||
# Очищаем имя от недопустимых символов
|
||
safe_name = "".join(c for c in project_name if c.isalnum() or c in "-_")
|
||
self.collection_name = safe_name
|
||
self._create_collection()
|
||
print(f"Переключено на проект: {self.collection_name}")
|
||
|
||
def get_current_project(self) -> str:
|
||
"""Получить имя текущего проекта"""
|
||
return self.collection_name
|
||
|
||
def list_projects(self) -> list:
|
||
"""Список всех проектов (коллекций)"""
|
||
try:
|
||
collections = self.client.get_collections()
|
||
return [c.name for c in collections.collections]
|
||
except:
|
||
return []
|
||
|
||
def delete_project(self, project_name: str = None):
|
||
"""Удалить проект (коллекцию)"""
|
||
name = project_name or self.collection_name
|
||
try:
|
||
self.client.delete_collection(collection_name=name)
|
||
print(f"Удалена коллекция: {name}")
|
||
except Exception as e:
|
||
print(f"Не удалось удалить коллекцию {name}: {e}")
|
||
|
||
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=500,
|
||
chunk_overlap=100,
|
||
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)
|
||
|
||
# Извлекаем название функции/класса/строки из контекста
|
||
metadata = doc.metadata
|
||
|
||
# Простая эвристика: ищем объявление функции/класса в начале чанка
|
||
import re
|
||
# Ищем паттерны: "void funcName(...)", "class ClassName", "int funcName(...)"
|
||
func_match = re.search(r'(?:class|struct|enum)\s+(\w+)', doc.page_content)
|
||
if not func_match:
|
||
func_match = re.search(r'(?:void|int|float|double|bool|auto)\s+(\w+)\s*\(', doc.page_content)
|
||
|
||
function_name = func_match.group(1) if func_match else "unknown"
|
||
|
||
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",
|
||
"function_name": function_name, # ✅ Добавил
|
||
"class_name": func_match.group(1) if func_match else None, # ✅ Добавил
|
||
"line_start": doc.metadata.get("line", 0), # Если есть в метаданных
|
||
},
|
||
# 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
|
||
|