Files
course-agent-od/backend/src/rag/vector_store.py
T
pengxiao 37c0364e7e refactor: replace print() with structured logging across backend
Replace all print()/stderr logging with Python logging module using
logger = logging.getLogger(__name__) pattern for consistent log levels
and formatting. Extract score conversion to shared score_utils module.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-05-28 15:52:44 +08:00

113 lines
3.1 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.
"""
LangChain 1.0 向量存储封装
"""
import os
import logging
import json
import hashlib
from typing import List, Dict, Any, Optional
from pathlib import Path
from langchain_chroma import Chroma
from langchain_core.documents import Document as LangChainDocument
from ..core.config import get_settings
from .embeddings import get_embedding_model
logger = logging.getLogger(__name__)
settings = get_settings()
class VectorStore:
"""向量存储管理器(LangChain 1.0 Chroma"""
def __init__(self):
"""初始化向量存储"""
self.vector_store_path = Path(settings.vector_store_path)
self.vector_store_path.mkdir(parents=True, exist_ok=True)
# 获取嵌入模型(保留自定义实现)
self.embedding_model = get_embedding_model()
# 使用LangChain Chroma wrapper
self.vectorstore = Chroma(
collection_name="course_knowledge",
embedding_function=self.embedding_model,
persist_directory=str(self.vector_store_path),
collection_metadata={"hnsw:space": "l2"} # 保持L2距离度量
)
def add_documents(self, documents: List[LangChainDocument]) -> bool:
"""添加LangChain文档到向量存储"""
try:
self.vectorstore.add_documents(documents)
return True
except Exception as e:
logger.error(f"添加文档失败: {str(e)}")
return False
def as_retriever(self, **kwargs):
"""返回标准LangChain检索器"""
return self.vectorstore.as_retriever(**kwargs)
def similarity_search_with_score(
self,
query: str,
k: int = 5,
filter: Optional[Dict] = None
):
"""相似度搜索(带分数)"""
logger.debug(f"查询参数 - k: {k}, filter: {filter}")
result = self.vectorstore.similarity_search_with_score(
query=query,
k=k,
filter=filter
)
logger.debug(f"返回结果数量: {len(result)}")
return result
def delete_by_document_id(self, document_id: int) -> bool:
"""删除指定文档的所有向量数据"""
try:
self.vectorstore._collection.delete(
where={"document_id": document_id}
)
return True
except Exception as e:
logger.error(f"删除向量数据失败: {str(e)}")
return False
def max_marginal_relevance_search(
self,
query: str,
k: int = 5,
fetch_k: int = 20,
filter: Optional[Dict] = None
):
"""MMR搜索(多样性检索)"""
return self.vectorstore.max_marginal_relevance_search(
query=query,
k=k,
fetch_k=fetch_k,
filter=filter
)
# 单例模式
_vector_store_instance = None
def get_vector_store() -> VectorStore:
"""获取向量存储实例"""
global _vector_store_instance
if _vector_store_instance is None:
_vector_store_instance = VectorStore()
return _vector_store_instance