fix(vector): harden packaged migrations
This commit is contained in:
@@ -27,7 +27,7 @@ COLLECTION_KINDS = {
|
||||
}
|
||||
ALLOWED_COLLECTIONS = frozenset(COLLECTION_KINDS)
|
||||
ALLOWED_KINDS = frozenset().union(*COLLECTION_KINDS.values())
|
||||
_VECTOR_DIMENSION = re.compile(r"^vector\((\d+)\)$")
|
||||
_VECTOR_DIMENSION = re.compile(r"^(?:[a-z_][a-z0-9_]*\.)?vector\((\d+)\)$")
|
||||
|
||||
|
||||
def _collection(schema: str, name: str) -> sql.Identifier:
|
||||
@@ -40,6 +40,14 @@ def _vector_literal(values: list[float]) -> str:
|
||||
return "[" + ",".join(str(float(value)) for value in values) + "]"
|
||||
|
||||
|
||||
def _vector_type(schema: str) -> sql.Identifier:
|
||||
return sql.Identifier(schema, "vector")
|
||||
|
||||
|
||||
def _cosine_operator(schema: str) -> sql.Composed:
|
||||
return sql.SQL("OPERATOR({}.<=>)").format(sql.Identifier(schema))
|
||||
|
||||
|
||||
def _validate_collection_kinds(collection: str, kinds: list[str]) -> None:
|
||||
invalid = set(kinds) - COLLECTION_KINDS[collection]
|
||||
if invalid:
|
||||
@@ -248,9 +256,16 @@ class PgVectorStore:
|
||||
continue
|
||||
where = sql.SQL(" WHERE kind = ANY(%s)") if collection_kinds else sql.SQL("")
|
||||
query = sql.SQL(
|
||||
"SELECT metadata, 1 - (embedding <=> %s::vector) AS similarity "
|
||||
"FROM {}{} ORDER BY embedding <=> %s::vector LIMIT %s"
|
||||
).format(table, where)
|
||||
"SELECT metadata, 1 - (embedding {} %s::{}) AS similarity "
|
||||
"FROM {}{} ORDER BY embedding {} %s::{} LIMIT %s"
|
||||
).format(
|
||||
_cosine_operator(self._schema),
|
||||
_vector_type(self._schema),
|
||||
table,
|
||||
where,
|
||||
_cosine_operator(self._schema),
|
||||
_vector_type(self._schema),
|
||||
)
|
||||
params = [_vector_literal(embedding)]
|
||||
if collection_kinds:
|
||||
params.append(collection_kinds)
|
||||
@@ -295,13 +310,13 @@ class PgVectorStore:
|
||||
raise VectorStoreError("Embedding dimension does not match configured dimension")
|
||||
insert = sql.SQL(
|
||||
"INSERT INTO {} (record_key, kind, content_hash, metadata, embedding) "
|
||||
"VALUES (%s, %s, %s, %s::jsonb, %s::vector) "
|
||||
"VALUES (%s, %s, %s, %s::jsonb, %s::{}) "
|
||||
"ON CONFLICT (record_key) DO NOTHING"
|
||||
).format(table)
|
||||
).format(table, _vector_type(self._schema))
|
||||
update = sql.SQL(
|
||||
"UPDATE {} SET kind = %s, content_hash = %s, metadata = %s::jsonb, "
|
||||
"embedding = %s::vector, indexed_at = now() WHERE record_key = %s"
|
||||
).format(table)
|
||||
"embedding = %s::{}, indexed_at = pg_catalog.now() WHERE record_key = %s"
|
||||
).format(table, _vector_type(self._schema))
|
||||
raw = engine.raw_connection()
|
||||
try:
|
||||
with raw.cursor() as cursor:
|
||||
|
||||
Reference in New Issue
Block a user