Files
NotesAgentic/backend/app/retrieval/vectorstore.py
yxxandClaude 6cf531f2a8 fix(retrieval): 修复原子性、过滤漏召回、tags 语义与 rebuild 回滚
- 元数据 + 向量单事务提交,避免 PATCH 半提交(审阅 #2)
- vectorstore upsert 改 delete-then-insert 幂等,支持共享 conn
- FTS 取全量 + 过滤 oversample,修复 metadata 过滤漏召回(审阅 #4)
- PATCH tags 区分 None/[]/非空:保留/清空/替换(审阅 #5)
- rebuild 拒绝增量 scope/note_ids,扫描先行 + 失败回滚旧索引(审阅 #6)

Co-Authored-By: Claude <noreply@anthropic.com>
2026-08-27 22:48:55 +08:00

95 lines
3.0 KiB
Python

"""VectorStore 统一接口与 sqlite-vec 实现。
vec0 虚拟表返回的 distance 是欧氏距离(非平方)。入库前向量已做 L2 归一化,
因此 distance² = 2(1-cos),余弦相似度 = 1 - distance² / 2。
"""
from __future__ import annotations
import sqlite3
from contextlib import nullcontext
from dataclasses import dataclass
from typing import Protocol, runtime_checkable
import sqlite_vec
from app.database.db import connect, transaction
@dataclass
class VectorRecord:
id: str
vector: list[float]
@dataclass
class VectorHit:
id: str
score: float # 余弦相似度 [0,1]
@runtime_checkable
class VectorStore(Protocol):
"""统一向量存储接口(与文档一致)。上层只依赖此抽象,不读 vec0 内部表。"""
async def upsert(self, records: list[VectorRecord]) -> None: ...
async def delete(self, ids: list[str]) -> None: ...
async def search(self, vector: list[float], *, top_k: int) -> list[VectorHit]: ...
class SqliteVecStore:
"""sqlite-vec 默认实现。"""
async def upsert(self, records: list[VectorRecord], *, conn: sqlite3.Connection | None = None) -> None:
if not records:
return
owns = conn is None
conn = conn or connect()
try:
with transaction(conn) if owns else nullcontext():
for record in records:
# vec0 不支持 UPDATE,采用 delete-then-insert 实现幂等 upsert,避免主键冲突
conn.execute("DELETE FROM vec_blocks WHERE block_id = ?", (record.id,))
conn.execute(
"INSERT INTO vec_blocks (block_id, embedding) VALUES (?, ?)",
(record.id, sqlite_vec.serialize_float32(record.vector)),
)
finally:
if owns:
conn.close()
async def delete(self, ids: list[str], *, conn: sqlite3.Connection | None = None) -> None:
if not ids:
return
owns = conn is None
conn = conn or connect()
try:
with transaction(conn) if owns else nullcontext():
for bid in ids:
conn.execute("DELETE FROM vec_blocks WHERE block_id = ?", (bid,))
finally:
if owns:
conn.close()
async def search(self, vector: list[float], *, top_k: int) -> list[VectorHit]:
conn = connect()
try:
rows = conn.execute(
"SELECT block_id, distance FROM vec_blocks WHERE embedding MATCH ? AND k = ?",
(sqlite_vec.serialize_float32(vector), top_k),
).fetchall()
return [
VectorHit(id=row["block_id"], score=max(0.0, 1.0 - row["distance"] ** 2 / 2.0))
for row in rows
]
finally:
conn.close()
async def clear(self) -> None:
conn = connect()
try:
with transaction(conn):
conn.execute("DELETE FROM vec_blocks")
finally:
conn.close()