Files
Agent/local_knowledge_base.py

204 lines
8.6 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, 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