37c0364e7e
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>
113 lines
3.1 KiB
Python
113 lines
3.1 KiB
Python
"""
|
||
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
|
||
|
||
|
||
|
||
|
||
|
||
|