nex_docus/backend/app/services/zvec_service.py

241 lines
6.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""
ZVec 向量化服务 - 调用阿里云ZVec API进行文档向量化
"""
import json
import os
import httpx
from typing import Dict, List, Optional, Any
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, delete
from app.models.document_vector import DocumentVector
from app.core.database import get_db
class ZVecService:
"""ZVec 向量化服务"""
ZVEC_API_ENDPOINT = os.getenv("ZVEC_API_ENDPOINT", "https://api.aliyun.com/zvec")
ZVEC_API_KEY = os.getenv("ZVEC_API_KEY", "")
ZVEC_MODEL = os.getenv("ZVEC_MODEL", "text-embedding-v3")
@classmethod
async def is_configured(cls) -> bool:
"""检查ZVec是否已配置"""
return bool(cls.ZVEC_API_KEY and cls.ZVEC_API_ENDPOINT)
@classmethod
async def vectorize_text(cls, text: str, doc_id: Optional[int] = None) -> Optional[List[float]]:
"""
调用ZVec API对文本进行向量化
Args:
text: 要向量化的文本内容
doc_id: 文档ID用于日志追踪
Returns:
向量列表或如果API调用失败则返回None
"""
if not await cls.is_configured():
return None
if not text or not text.strip():
return None
try:
# 截断超长文本ZVec有输入限制
text = text[:8000]
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {cls.ZVEC_API_KEY}",
}
payload = {
"model": cls.ZVEC_MODEL,
"input": text,
}
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.post(
f"{cls.ZVEC_API_ENDPOINT}/embeddings",
headers=headers,
json=payload,
)
response.raise_for_status()
data = response.json()
embeddings = data.get("data", [])
if embeddings:
return embeddings[0].get("embedding", [])
return None
except Exception as exc:
print(f"ZVec vectorization failed for doc_id={doc_id}: {exc}")
return None
@classmethod
async def store_vector(
cls,
db: AsyncSession,
project_id: int,
file_path: str,
content: str,
doc_id: Optional[int] = None,
) -> bool:
"""
向量化文本并存储到数据库
Args:
db: 数据库会话
project_id: 项目ID
file_path: 文件路径
content: 文件内容
doc_id: 文档元数据ID
Returns:
是否成功存储
"""
# 调用ZVec获取向量
embedding = await cls.vectorize_text(content, doc_id)
if not embedding:
return False
try:
# 先删除旧的向量记录
await db.execute(
delete(DocumentVector).where(
(DocumentVector.project_id == project_id)
& (DocumentVector.file_path == file_path)
)
)
# 创建新的向量记录
vector_record = DocumentVector(
project_id=project_id,
file_path=file_path,
doc_id=doc_id,
embedding=json.dumps(embedding),
embedding_model=cls.ZVEC_MODEL,
chunk_count=1, # 单个文件作为一个chunk
)
db.add(vector_record)
await db.commit()
return True
except Exception as exc:
print(f"Failed to store vector for {file_path}: {exc}")
await db.rollback()
return False
@classmethod
async def delete_vector(
cls,
db: AsyncSession,
project_id: int,
file_path: str,
) -> bool:
"""
删除文档的向量记录
Args:
db: 数据库会话
project_id: 项目ID
file_path: 文件路径
Returns:
是否成功删除
"""
try:
await db.execute(
delete(DocumentVector).where(
(DocumentVector.project_id == project_id)
& (DocumentVector.file_path == file_path)
)
)
await db.commit()
return True
except Exception as exc:
print(f"Failed to delete vector for {file_path}: {exc}")
await db.rollback()
return False
@classmethod
async def search_similar(
cls,
db: AsyncSession,
project_id: int,
query: str,
top_k: int = 5,
) -> List[Dict[str, Any]]:
"""
基于查询文本检索相似文档
Args:
db: 数据库会话
project_id: 项目ID
query: 查询文本
top_k: 返回最相似的文档数量
Returns:
相似文档列表,按相似度排序
"""
# 获取查询文本的向量
query_embedding = await cls.vectorize_text(query)
if not query_embedding:
return []
try:
# 获取项目中所有文档的向量
result = await db.execute(
select(DocumentVector).where(
DocumentVector.project_id == project_id
)
)
vectors = result.scalars().all()
if not vectors:
return []
# 计算相似度(余弦相似度)
similarities = []
for vector_record in vectors:
try:
doc_embedding = json.loads(vector_record.embedding)
similarity = cls._cosine_similarity(query_embedding, doc_embedding)
similarities.append({
"file_path": vector_record.file_path,
"doc_id": vector_record.doc_id,
"similarity": similarity,
})
except json.JSONDecodeError:
continue
# 按相似度排序并返回top-k
similarities.sort(key=lambda x: x["similarity"], reverse=True)
return similarities[:top_k]
except Exception as exc:
print(f"Failed to search similar documents: {exc}")
return []
@staticmethod
def _cosine_similarity(vec1: List[float], vec2: List[float]) -> float:
"""计算两个向量的余弦相似度"""
if not vec1 or not vec2 or len(vec1) != len(vec2):
return 0.0
dot_product = sum(a * b for a, b in zip(vec1, vec2))
magnitude1 = sum(a * a for a in vec1) ** 0.5
magnitude2 = sum(b * b for b in vec2) ** 0.5
if magnitude1 == 0 or magnitude2 == 0:
return 0.0
return dot_product / (magnitude1 * magnitude2)
zvec_service = ZVecService()