241 lines
6.9 KiB
Python
241 lines
6.9 KiB
Python
|
|
"""
|
|||
|
|
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()
|