717 lines
25 KiB
Python
717 lines
25 KiB
Python
"""
|
||
ZVec 本地向量化服务 - 使用本地 zvec 库 + 模型配置中的 embedding 模型
|
||
"""
|
||
import os
|
||
import hashlib
|
||
import logging
|
||
from pathlib import Path
|
||
from typing import Dict, List, Optional, Any
|
||
|
||
import zvec
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
from sqlalchemy import select, delete
|
||
|
||
from app.models.llm_model_config import LLMModelConfig
|
||
from app.models.document_vector import DocumentVector
|
||
from app.services.llm_provider_service import LLMProviderService
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
def _resolve_zvec_data_dir() -> str:
|
||
"""解析 ZVec 向量库根目录。
|
||
|
||
优先用 ZVEC_DATA_DIR 环境变量;否则放到 storage 存储区下的
|
||
vector_index 目录,与 search_index 并列,便于统一备份/管理。
|
||
"""
|
||
env_dir = os.getenv("ZVEC_DATA_DIR")
|
||
if env_dir:
|
||
return env_dir
|
||
|
||
from app.core.config import settings
|
||
|
||
storage_root = Path(settings.STORAGE_ROOT)
|
||
if not storage_root.is_absolute():
|
||
backend_dir = Path(__file__).parent.parent.parent
|
||
storage_root = (backend_dir / storage_root).resolve()
|
||
return str(storage_root / "vector_index")
|
||
|
||
|
||
ZVEC_DATA_DIR = _resolve_zvec_data_dir()
|
||
EMBEDDING_DIMENSION = int(os.getenv("ZVEC_EMBEDDING_DIM", "1536"))
|
||
|
||
|
||
class ZVecService:
|
||
"""基于本地 zvec 库的向量存储和检索服务"""
|
||
|
||
_collections: Dict[int, zvec.Collection] = {}
|
||
|
||
@classmethod
|
||
def _get_collection_path(cls, project_id: int) -> str:
|
||
"""返回项目 collection 路径。
|
||
|
||
只确保父目录存在;叶子目录由 zvec.create_and_open 自行创建,
|
||
预先创建叶子目录会让 create_and_open 报 "path exists"。
|
||
"""
|
||
path = Path(ZVEC_DATA_DIR) / str(project_id)
|
||
path.parent.mkdir(parents=True, exist_ok=True)
|
||
return str(path)
|
||
|
||
@classmethod
|
||
def _open_collection(cls, project_id: int) -> Optional[zvec.Collection]:
|
||
"""打开已存在的 collection;不存在则返回 None"""
|
||
if project_id in cls._collections:
|
||
return cls._collections[project_id]
|
||
|
||
collection_path = cls._get_collection_path(project_id)
|
||
try:
|
||
collection = zvec.open(path=collection_path)
|
||
except Exception:
|
||
return None
|
||
|
||
cls._collections[project_id] = collection
|
||
return collection
|
||
|
||
@classmethod
|
||
def _get_or_create_collection(cls, project_id: int, dimension: int) -> zvec.Collection:
|
||
"""按指定维度获取或创建 collection。
|
||
|
||
ZVec 的 collection 维度创建后固定,若已存在的维度与当前 embedding
|
||
模型不一致,则重建该项目的 collection(旧向量作废,需重新向量化)。
|
||
"""
|
||
collection = cls._open_collection(project_id)
|
||
if collection is not None:
|
||
existing_dim = cls._collection_dimension(collection)
|
||
if existing_dim is not None and existing_dim != dimension:
|
||
logger.warning(
|
||
"Project %s collection dim %s != model dim %s, rebuilding",
|
||
project_id, existing_dim, dimension,
|
||
)
|
||
cls._drop_collection(project_id)
|
||
collection = None
|
||
else:
|
||
return collection
|
||
|
||
collection_path = cls._get_collection_path(project_id)
|
||
# ZVec 的 create_and_open 要求目标路径不存在,清掉可能残留的空/损坏目录,保证幂等
|
||
leaf = Path(ZVEC_DATA_DIR) / str(project_id)
|
||
if leaf.exists():
|
||
import shutil
|
||
shutil.rmtree(leaf, ignore_errors=True)
|
||
schema = zvec.CollectionSchema(
|
||
name=f"project_{project_id}",
|
||
vectors=zvec.VectorSchema(
|
||
"embedding", zvec.DataType.VECTOR_FP32, dimension
|
||
),
|
||
)
|
||
collection = zvec.create_and_open(path=collection_path, schema=schema)
|
||
cls._collections[project_id] = collection
|
||
return collection
|
||
|
||
@staticmethod
|
||
def _collection_dimension(collection: zvec.Collection) -> Optional[int]:
|
||
"""尽力读取 collection 的向量维度,读取失败返回 None"""
|
||
try:
|
||
schema = collection.schema
|
||
vectors = getattr(schema, "vectors", None)
|
||
if vectors is None:
|
||
return None
|
||
vector_schema = vectors[0] if isinstance(vectors, (list, tuple)) else vectors
|
||
return getattr(vector_schema, "dimension", None)
|
||
except Exception:
|
||
return None
|
||
|
||
@classmethod
|
||
def _drop_collection(cls, project_id: int) -> None:
|
||
"""删除项目 collection 的本地数据并清理缓存"""
|
||
cls._collections.pop(project_id, None)
|
||
import shutil
|
||
|
||
path = Path(ZVEC_DATA_DIR) / str(project_id)
|
||
if path.exists():
|
||
shutil.rmtree(path, ignore_errors=True)
|
||
|
||
@classmethod
|
||
async def get_embedding_config(cls, db: AsyncSession) -> Optional[LLMModelConfig]:
|
||
"""获取系统配置的 embedding 模型(按 model_type 精确查询)"""
|
||
stmt = (
|
||
select(LLMModelConfig)
|
||
.where(
|
||
LLMModelConfig.is_active == True,
|
||
LLMModelConfig.model_type == "embedding",
|
||
)
|
||
.order_by(
|
||
LLMModelConfig.is_default.desc(),
|
||
LLMModelConfig.updated_at.desc(),
|
||
LLMModelConfig.config_id.desc(),
|
||
)
|
||
)
|
||
result = await db.execute(stmt)
|
||
return result.scalars().first()
|
||
|
||
@classmethod
|
||
async def generate_embedding(
|
||
cls, db: AsyncSession, text: str
|
||
) -> Optional[List[float]]:
|
||
"""调用配置的 embedding 模型生成向量"""
|
||
if not text or not text.strip():
|
||
return None
|
||
|
||
config = await cls.get_embedding_config(db)
|
||
if not config:
|
||
logger.warning("No embedding model configured, skipping vectorization")
|
||
return None
|
||
|
||
try:
|
||
return await LLMProviderService.generate_embedding(
|
||
provider=config.provider,
|
||
endpoint_url=config.endpoint_url,
|
||
api_key=config.api_key,
|
||
llm_model_name=config.llm_model_name,
|
||
text=text,
|
||
timeout=config.llm_timeout or 60,
|
||
dimension=config.embedding_dimension,
|
||
)
|
||
except Exception as exc:
|
||
logger.error(f"Embedding generation failed: {exc}")
|
||
return None
|
||
|
||
@staticmethod
|
||
def _doc_id(file_path: str, chunk_index: int = 0) -> str:
|
||
"""生成 ZVec 合法的 doc id。
|
||
|
||
ZVec 的 doc id 不允许中文等非 ASCII 字符,故对文件路径做 MD5 哈希。
|
||
分块后每个 chunk 需独立 doc id,故追加 chunk 序号。
|
||
file_path 的反查通过 document_vector 表的 zvec_id 字段完成。
|
||
"""
|
||
base = hashlib.md5(file_path.encode("utf-8")).hexdigest()
|
||
return f"{base}#{chunk_index}"
|
||
|
||
@staticmethod
|
||
def _content_hash(content: str) -> str:
|
||
return hashlib.sha256((content or "").encode("utf-8")).hexdigest()
|
||
|
||
@staticmethod
|
||
def _make_anchor(text: str, max_chars: int = 120) -> str:
|
||
"""从 chunk 文本中提取定位锚点:首个非空文本行、去除常见 markdown 标记。
|
||
|
||
锚点供旧引用定位兼容使用,故取纯文本片段。
|
||
"""
|
||
import re
|
||
|
||
for raw_line in (text or "").splitlines():
|
||
line = raw_line.strip()
|
||
if not line:
|
||
continue
|
||
# 去除标题井号、列表符号、引用符号等行首标记
|
||
line = re.sub(r"^\s*#{1,6}\s+", "", line)
|
||
line = re.sub(r"^\s*[-*+]\s+", "", line)
|
||
line = re.sub(r"^\s*>\s+", "", line)
|
||
line = re.sub(r"^\s*\d+\.\s+", "", line)
|
||
# 去除行内强调/代码标记
|
||
line = re.sub(r"[*_`~]", "", line)
|
||
line = line.strip()
|
||
if line:
|
||
return line[:max_chars]
|
||
return (text or "").strip()[:max_chars]
|
||
|
||
@classmethod
|
||
def _chunk_text(
|
||
cls,
|
||
content: str,
|
||
chunk_size: int,
|
||
overlap: int,
|
||
) -> List[Dict[str, Any]]:
|
||
"""按字符滑动窗口分块。
|
||
|
||
返回 [{index, text, anchor}]。相邻分块重叠 overlap 个字符,
|
||
以避免语义在分块边界被切断。空内容返回空列表。
|
||
"""
|
||
text = content or ""
|
||
if not text.strip():
|
||
return []
|
||
|
||
chunk_size = max(1, int(chunk_size))
|
||
overlap = max(0, min(int(overlap), chunk_size - 1))
|
||
step = chunk_size - overlap
|
||
|
||
chunks: List[Dict[str, Any]] = []
|
||
start = 0
|
||
length = len(text)
|
||
index = 0
|
||
while start < length:
|
||
piece = text[start:start + chunk_size]
|
||
if piece.strip():
|
||
chunks.append({
|
||
"index": index,
|
||
"text": piece,
|
||
"anchor": cls._make_anchor(piece),
|
||
})
|
||
index += 1
|
||
start += step
|
||
|
||
return chunks
|
||
|
||
@classmethod
|
||
async def _upsert_vector_record(
|
||
cls,
|
||
db: AsyncSession,
|
||
project_id: int,
|
||
file_path: str,
|
||
*,
|
||
status: str,
|
||
chunk_index: int = 0,
|
||
chunk_text: Optional[str] = None,
|
||
content_hash: Optional[str] = None,
|
||
zvec_id: Optional[str] = None,
|
||
error_message: Optional[str] = None,
|
||
) -> None:
|
||
"""写入/更新 document_vector 表中某个分块的向量化状态。
|
||
|
||
按 (project_id, file_path, chunk_index) 定位记录。失败状态通常写
|
||
chunk_index=0 的一条即可,成功状态则逐分块写入。
|
||
"""
|
||
stmt = select(DocumentVector).where(
|
||
DocumentVector.project_id == project_id,
|
||
DocumentVector.file_path == file_path,
|
||
DocumentVector.chunk_index == chunk_index,
|
||
)
|
||
result = await db.execute(stmt)
|
||
record = result.scalar_one_or_none()
|
||
|
||
if record is None:
|
||
record = DocumentVector(
|
||
project_id=project_id,
|
||
file_path=file_path,
|
||
chunk_index=chunk_index,
|
||
chunk_text=chunk_text,
|
||
status=status,
|
||
content_hash=content_hash,
|
||
zvec_id=zvec_id,
|
||
error_message=(error_message or "")[:500] or None,
|
||
)
|
||
db.add(record)
|
||
else:
|
||
record.status = status
|
||
if chunk_text is not None:
|
||
record.chunk_text = chunk_text
|
||
if content_hash is not None:
|
||
record.content_hash = content_hash
|
||
if zvec_id is not None:
|
||
record.zvec_id = zvec_id
|
||
record.error_message = (error_message or "")[:500] or None
|
||
|
||
@classmethod
|
||
async def _purge_file_chunks(
|
||
cls,
|
||
db: AsyncSession,
|
||
project_id: int,
|
||
file_path: str,
|
||
) -> None:
|
||
"""删除某文件在 ZVec 与 document_vector 表中的所有分块记录(不 commit)。"""
|
||
# 先查出该文件所有分块的 zvec_id,用于从 ZVec 集合中删除
|
||
result = await db.execute(
|
||
select(DocumentVector.zvec_id).where(
|
||
DocumentVector.project_id == project_id,
|
||
DocumentVector.file_path == file_path,
|
||
)
|
||
)
|
||
zvec_ids = [z for (z,) in result.all() if z]
|
||
if zvec_ids:
|
||
try:
|
||
collection = cls._open_collection(project_id)
|
||
if collection is not None:
|
||
collection.delete(zvec_ids)
|
||
except Exception as exc:
|
||
logger.error(f"ZVec delete chunks failed for {file_path}: {exc}")
|
||
|
||
await db.execute(
|
||
delete(DocumentVector).where(
|
||
DocumentVector.project_id == project_id,
|
||
DocumentVector.file_path == file_path,
|
||
)
|
||
)
|
||
|
||
@classmethod
|
||
async def vectorize_markdown(
|
||
cls,
|
||
db: AsyncSession,
|
||
project_id: int,
|
||
file_path: str,
|
||
content: str,
|
||
) -> bool:
|
||
"""对 MD 文件分块向量化并存入 ZVec,同时逐分块记录状态到 document_vector。
|
||
|
||
流程:删除旧分块 → 分块 → 逐块生成 embedding、写入 ZVec、写记录。
|
||
"""
|
||
from app.core.config import settings
|
||
|
||
content_hash = cls._content_hash(content)
|
||
|
||
# 先清理该文件的所有旧分块(增量重建),保证不残留过期向量
|
||
await cls._purge_file_chunks(db, project_id, file_path)
|
||
|
||
chunks = cls._chunk_text(content, settings.CHUNK_SIZE, settings.CHUNK_OVERLAP)
|
||
if not chunks:
|
||
# 空文件:清理后直接提交,视为成功(无可向量化内容)
|
||
await db.commit()
|
||
return True
|
||
|
||
success_count = 0
|
||
last_error: Optional[str] = None
|
||
for chunk in chunks:
|
||
embedding = await cls.generate_embedding(db, chunk["text"])
|
||
if not embedding:
|
||
last_error = "未配置可用的 embedding 模型或向量生成失败"
|
||
continue
|
||
|
||
doc_id = cls._doc_id(file_path, chunk["index"])
|
||
try:
|
||
collection = cls._get_or_create_collection(project_id, len(embedding))
|
||
collection.insert([
|
||
zvec.Doc(id=doc_id, vectors={"embedding": embedding})
|
||
])
|
||
except Exception as exc:
|
||
logger.error(
|
||
f"ZVec insert failed for {file_path} chunk {chunk['index']}: {exc}"
|
||
)
|
||
last_error = str(exc)
|
||
continue
|
||
|
||
await cls._upsert_vector_record(
|
||
db, project_id, file_path,
|
||
status="success",
|
||
chunk_index=chunk["index"],
|
||
chunk_text=chunk["text"],
|
||
content_hash=content_hash,
|
||
zvec_id=doc_id,
|
||
)
|
||
success_count += 1
|
||
|
||
if success_count == 0:
|
||
# 全部分块失败:写一条 failed 记录(chunk_index=0)便于进度展示
|
||
await cls._upsert_vector_record(
|
||
db, project_id, file_path,
|
||
status="failed",
|
||
chunk_index=0,
|
||
content_hash=content_hash,
|
||
error_message=last_error or "向量化失败",
|
||
)
|
||
await db.commit()
|
||
return False
|
||
|
||
await db.commit()
|
||
return True
|
||
|
||
@classmethod
|
||
async def delete_vector(
|
||
cls,
|
||
db: AsyncSession,
|
||
project_id: int,
|
||
file_path: str,
|
||
) -> bool:
|
||
"""从 ZVec 删除文档的所有分块向量,并清除 document_vector 记录"""
|
||
try:
|
||
await cls._purge_file_chunks(db, project_id, file_path)
|
||
await db.commit()
|
||
except Exception as exc:
|
||
logger.error(f"Delete document_vector record failed for {file_path}: {exc}")
|
||
await db.rollback()
|
||
return True
|
||
|
||
# 兼容 project_file_service 中的调用别名
|
||
@classmethod
|
||
async def delete_vectors(
|
||
cls,
|
||
db: AsyncSession,
|
||
project_id: int,
|
||
file_path: str,
|
||
) -> bool:
|
||
"""delete_vector 的别名(兼容文件服务调用)"""
|
||
return await cls.delete_vector(db, project_id, file_path)
|
||
|
||
@classmethod
|
||
async def sync_vector_move(
|
||
cls,
|
||
db: AsyncSession,
|
||
project_id: int,
|
||
old_path: str,
|
||
new_path: str,
|
||
content: Optional[str] = None,
|
||
) -> bool:
|
||
"""文件移动/重命名时同步向量。
|
||
|
||
content 为空时会尝试读取新路径文件内容再向量化。
|
||
"""
|
||
await cls.delete_vector(db, project_id, old_path)
|
||
if content is None:
|
||
return await cls.revectorize_path(db, project_id, new_path)
|
||
return await cls.vectorize_markdown(db, project_id, new_path, content)
|
||
|
||
@classmethod
|
||
async def update_vector_path(
|
||
cls,
|
||
db: AsyncSession,
|
||
project_id: int,
|
||
old_path: str,
|
||
new_path: str,
|
||
content: str,
|
||
) -> bool:
|
||
"""文件重命名/移动时更新向量(删除旧的,插入新的)"""
|
||
await cls.delete_vector(db, project_id, old_path)
|
||
return await cls.vectorize_markdown(db, project_id, new_path, content)
|
||
|
||
@classmethod
|
||
async def revectorize_path(
|
||
cls,
|
||
db: AsyncSession,
|
||
project_id: int,
|
||
file_path: str,
|
||
) -> bool:
|
||
"""读取项目中指定文件内容并重新向量化"""
|
||
content = await cls._read_project_file(db, project_id, file_path)
|
||
if content is None:
|
||
return False
|
||
return await cls.vectorize_markdown(db, project_id, file_path, content)
|
||
|
||
@staticmethod
|
||
async def _read_project_file(
|
||
db: AsyncSession,
|
||
project_id: int,
|
||
file_path: str,
|
||
) -> Optional[str]:
|
||
"""读取项目内某个相对路径文件的文本内容"""
|
||
from app.models.project import Project
|
||
from app.services.storage import storage_service
|
||
|
||
result = await db.execute(select(Project).where(Project.id == project_id))
|
||
project = result.scalar_one_or_none()
|
||
if not project:
|
||
return None
|
||
try:
|
||
full_path = storage_service.get_secure_path(project.storage_key, file_path)
|
||
return await storage_service.read_file(full_path)
|
||
except Exception as exc:
|
||
logger.error(f"Read project file failed {file_path}: {exc}")
|
||
return None
|
||
|
||
@classmethod
|
||
def _list_markdown_files(cls, storage_key: str) -> List[str]:
|
||
"""列出项目下所有 MD 文件的相对路径(排除 _assets 资源目录)"""
|
||
from app.services.storage import storage_service
|
||
|
||
root = storage_service.get_secure_path(storage_key)
|
||
if not root.exists():
|
||
return []
|
||
files: List[str] = []
|
||
for md_path in root.rglob("*.md"):
|
||
try:
|
||
rel = md_path.relative_to(root).as_posix()
|
||
except ValueError:
|
||
continue
|
||
# 跳过 _assets 资源目录
|
||
if rel.startswith("_assets/") or "/_assets/" in rel:
|
||
continue
|
||
# 跳过 . 开头的隐藏文件或位于隐藏目录下的文件(如 .git/.obsidian 等)
|
||
if any(part.startswith(".") for part in rel.split("/")):
|
||
continue
|
||
files.append(rel)
|
||
return files
|
||
|
||
@staticmethod
|
||
def _aggregate_file_records(
|
||
rows: List[DocumentVector],
|
||
) -> Dict[str, Dict[str, Any]]:
|
||
"""将多分块记录按文件聚合为文件级状态。
|
||
|
||
返回 {file_path: {status, content_hash, error_message}}。
|
||
任一分块 failed 则该文件视为 failed;否则若有 success 则 success。
|
||
"""
|
||
by_file: Dict[str, Dict[str, Any]] = {}
|
||
for r in rows:
|
||
agg = by_file.get(r.file_path)
|
||
if agg is None:
|
||
agg = {"status": r.status, "content_hash": r.content_hash, "error_message": r.error_message}
|
||
by_file[r.file_path] = agg
|
||
else:
|
||
# failed 优先级最高,用于暴露问题
|
||
if r.status == "failed":
|
||
agg["status"] = "failed"
|
||
agg["error_message"] = r.error_message or agg.get("error_message")
|
||
elif agg["status"] != "failed" and r.status == "success":
|
||
agg["status"] = "success"
|
||
if agg.get("content_hash") is None and r.content_hash:
|
||
agg["content_hash"] = r.content_hash
|
||
return by_file
|
||
|
||
@classmethod
|
||
async def get_progress(
|
||
cls,
|
||
db: AsyncSession,
|
||
project_id: int,
|
||
storage_key: str,
|
||
) -> Dict[str, Any]:
|
||
"""统计项目向量化进度:已成功 / 失败 / 待处理 / 总 MD 文件数"""
|
||
md_files = cls._list_markdown_files(storage_key)
|
||
total = len(md_files)
|
||
|
||
result = await db.execute(
|
||
select(DocumentVector).where(DocumentVector.project_id == project_id)
|
||
)
|
||
records = cls._aggregate_file_records(list(result.scalars().all()))
|
||
|
||
success = 0
|
||
failed_items: List[Dict[str, str]] = []
|
||
pending_items: List[str] = []
|
||
for path in md_files:
|
||
record = records.get(path)
|
||
if record and record["status"] == "success":
|
||
success += 1
|
||
elif record and record["status"] == "failed":
|
||
failed_items.append({
|
||
"file_path": path,
|
||
"error": record.get("error_message") or "向量化失败",
|
||
})
|
||
else:
|
||
pending_items.append(path)
|
||
|
||
percent = int(success / total * 100) if total else 0
|
||
return {
|
||
"total": total,
|
||
"success": success,
|
||
"failed": len(failed_items),
|
||
"pending": len(pending_items),
|
||
"percent": percent,
|
||
"failed_items": failed_items[:50],
|
||
"pending_items": pending_items[:50],
|
||
"embedding_ready": await cls.get_embedding_config(db) is not None,
|
||
}
|
||
|
||
@classmethod
|
||
async def vectorize_project(
|
||
cls,
|
||
db: AsyncSession,
|
||
project_id: int,
|
||
storage_key: str,
|
||
*,
|
||
force: bool = False,
|
||
) -> Dict[str, Any]:
|
||
"""批量向量化项目下所有 MD 文件。
|
||
|
||
force=False 时跳过内容未变化且已成功的文件(增量);
|
||
force=True 时全部重新向量化(全量)。
|
||
"""
|
||
md_files = cls._list_markdown_files(storage_key)
|
||
md_file_set = set(md_files)
|
||
|
||
result = await db.execute(
|
||
select(DocumentVector).where(DocumentVector.project_id == project_id)
|
||
)
|
||
records = cls._aggregate_file_records(list(result.scalars().all()))
|
||
|
||
if force:
|
||
cls._drop_collection(project_id)
|
||
await db.execute(
|
||
delete(DocumentVector).where(DocumentVector.project_id == project_id)
|
||
)
|
||
await db.commit()
|
||
records = {}
|
||
else:
|
||
for stale_path in set(records) - md_file_set:
|
||
await cls.delete_vector(db, project_id, stale_path)
|
||
records.pop(stale_path, None)
|
||
|
||
processed = 0
|
||
skipped = 0
|
||
failed = 0
|
||
for path in md_files:
|
||
content = await cls._read_project_file(db, project_id, path)
|
||
if content is None:
|
||
failed += 1
|
||
continue
|
||
|
||
if not force:
|
||
record = records.get(path)
|
||
if (
|
||
record
|
||
and record["status"] == "success"
|
||
and record.get("content_hash") == cls._content_hash(content)
|
||
):
|
||
skipped += 1
|
||
continue
|
||
|
||
ok = await cls.vectorize_markdown(db, project_id, path, content)
|
||
if ok:
|
||
processed += 1
|
||
else:
|
||
failed += 1
|
||
|
||
return {
|
||
"total": len(md_files),
|
||
"processed": processed,
|
||
"skipped": skipped,
|
||
"failed": failed,
|
||
}
|
||
|
||
@classmethod
|
||
async def search_similar(
|
||
cls,
|
||
db: AsyncSession,
|
||
project_id: int,
|
||
query: str,
|
||
top_k: int = 5,
|
||
) -> List[Dict[str, Any]]:
|
||
"""基于查询文本检索相似文档"""
|
||
query_embedding = await cls.generate_embedding(db, query)
|
||
if not query_embedding:
|
||
return []
|
||
|
||
try:
|
||
collection = cls._open_collection(project_id)
|
||
if collection is None:
|
||
return [] # 项目尚未向量化
|
||
results = collection.query(
|
||
zvec.VectorQuery("embedding", vector=query_embedding),
|
||
topk=top_k,
|
||
)
|
||
|
||
# doc id 是「路径哈希#分块序号」,通过 document_vector 表的 zvec_id
|
||
# 反查真实 file_path、chunk_index 及分块文本 chunk_text
|
||
zvec_ids = [r.id for r in results]
|
||
if not zvec_ids:
|
||
return []
|
||
|
||
db_result = await db.execute(
|
||
select(DocumentVector).where(
|
||
DocumentVector.project_id == project_id,
|
||
DocumentVector.zvec_id.in_(zvec_ids),
|
||
)
|
||
)
|
||
record_by_id = {r.zvec_id: r for r in db_result.scalars().all()}
|
||
|
||
matched_docs = []
|
||
for result in results:
|
||
record = record_by_id.get(result.id)
|
||
if not record:
|
||
continue
|
||
matched_docs.append({
|
||
"file_path": record.file_path,
|
||
"chunk_index": record.chunk_index,
|
||
"chunk_text": record.chunk_text or "",
|
||
"score": result.score if hasattr(result, "score") else 0.0,
|
||
})
|
||
|
||
return matched_docs
|
||
|
||
except Exception as exc:
|
||
logger.error(f"ZVec search failed: {exc}")
|
||
return []
|
||
|
||
@classmethod
|
||
def close_collection(cls, project_id: int) -> None:
|
||
"""关闭项目的 ZVec collection"""
|
||
if project_id in cls._collections:
|
||
del cls._collections[project_id]
|
||
|
||
|
||
zvec_service = ZVecService()
|