Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion engine/clients/pgvector/configure.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,10 +24,14 @@ def clean(self):
def recreate(self, dataset: Dataset, collection_params):
if dataset.config.distance == Distance.DOT:
raise IncompatibilityError
if "geo" in dataset.config.schema.values():
# payload is a flat JSONB blob, no geo radius search; skip before upload
raise IncompatibilityError

self.conn.execute(f"""CREATE TABLE items (
id SERIAL PRIMARY KEY,
embedding vector({dataset.config.vector_size}) NOT NULL
embedding vector({dataset.config.vector_size}) NOT NULL,
payload JSONB
);""")
self.conn.execute("ALTER TABLE items ALTER COLUMN embedding SET STORAGE PLAIN")

Expand Down
78 changes: 56 additions & 22 deletions engine/clients/pgvector/parser.py
Original file line number Diff line number Diff line change
@@ -1,24 +1,57 @@
import json
from typing import Any, List, Optional
from typing import Any, Dict, List, Optional, Tuple

from psycopg.types.json import Jsonb

from engine.base_client import IncompatibilityError
from engine.base_client.parser import BaseConditionParser, FieldValue

# (sql_clause, params) — clause references payload fields only through bound
# parameters, never string-interpolated, so values never need manual SQL quoting.
ParsedCondition = Tuple[str, Dict[str, Any]]


class PgVectorConditionParser(BaseConditionParser):
def __init__(self) -> None:
super().__init__()
self._counter = 0

def _next_param(self) -> str:
self._counter += 1
return f"p{self._counter}"

def build_condition(
self, and_subfilters: Optional[List[Any]], or_subfilters: Optional[List[Any]]
) -> Optional[Any]:
self,
and_subfilters: Optional[List[ParsedCondition]],
or_subfilters: Optional[List[ParsedCondition]],
) -> Optional[ParsedCondition]:
clauses = []
if or_subfilters is not None and len(or_subfilters) > 0:
clauses.append(f"( {' OR '.join(or_subfilters)} )")
if and_subfilters is not None and len(and_subfilters) > 0:
clauses.append(f"( {' AND '.join(and_subfilters)} )")
params: Dict[str, Any] = {}

if or_subfilters:
sub_clauses, sub_params = zip(*or_subfilters)
clauses.append("( " + " OR ".join(sub_clauses) + " )")
for p in sub_params:
params.update(p)
if and_subfilters:
sub_clauses, sub_params = zip(*and_subfilters)
clauses.append("( " + " AND ".join(sub_clauses) + " )")
for p in sub_params:
params.update(p)

return " AND ".join(clauses)
return " AND ".join(clauses), params

def build_exact_match_filter(self, field_name: str, value: FieldValue) -> Any:
return f"{field_name} == {json.dumps(value)}"
def build_exact_match_filter(
self, field_name: str, value: FieldValue
) -> ParsedCondition:
key_param = self._next_param()
value_param = self._next_param()
# `@>` not `=`: equality for scalar fields, "contains" for list-valued
# fields (e.g. arxiv `labels`), same semantics as Qdrant MatchValue and
# Elasticsearch `match` on an array.
return (
f"payload -> %({key_param})s @> %({value_param})s",
{key_param: field_name, value_param: Jsonb(value)},
)

def build_range_filter(
self,
Expand All @@ -27,20 +60,21 @@ def build_range_filter(
gt: Optional[FieldValue],
lte: Optional[FieldValue],
gte: Optional[FieldValue],
) -> Any:
) -> ParsedCondition:
clauses = []
if lt is not None:
clauses.append(f"{field_name} < {lt}")
if gt is not None:
clauses.append(f"{field_name} > {gt}")
if lte is not None:
clauses.append(f"{field_name} <= {lte}")
if gte is not None:
clauses.append(f"{field_name} >= {gte}")
return f"( {' AND '.join(clauses)} )"
params: Dict[str, Any] = {}
for op, bound in (("<", lt), (">", gt), ("<=", lte), (">=", gte)):
if bound is None:
continue
key_param = self._next_param()
value_param = self._next_param()
clauses.append(f"payload -> %({key_param})s {op} %({value_param})s")
params[key_param] = field_name
params[value_param] = Jsonb(bound)
return "( " + " AND ".join(clauses) + " )", params

def build_geo_filter(
self, field_name: str, lat: float, lon: float, radius: float
) -> Any:
# TODO: Implement this
# payload is a flat JSONB blob, no geo type/indexing support
raise IncompatibilityError
16 changes: 12 additions & 4 deletions engine/clients/pgvector/search.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,17 +25,25 @@ def init_client(cls, host, distance, connection_params: dict, search_params: dic
cls.cur = cls.conn.cursor()
cls.cur.execute(f"SET hnsw.ef_search = {search_params['config']['hnsw_ef']}")
if distance == Distance.COSINE:
cls.query = "SELECT id, embedding <=> %s AS _score FROM items ORDER BY _score LIMIT %s"
cls.query_template = "SELECT id, embedding <=> %(vector)s AS _score FROM items{where} ORDER BY _score LIMIT %(top)s"
elif distance == Distance.L2:
cls.query = "SELECT id, embedding <-> %s AS _score FROM items ORDER BY _score LIMIT %s"
cls.query_template = "SELECT id, embedding <-> %(vector)s AS _score FROM items{where} ORDER BY _score LIMIT %(top)s"
else:
raise NotImplementedError(f"Unsupported distance metric {cls.distance}")

@classmethod
def search_one(cls, query: Query, top) -> List[Tuple[int, float]]:
# TODO: Use query.metaconditions for datasets with filtering
params = {"vector": np.array(query.vector), "top": top}
condition = cls.parser.parse(query.meta_conditions)
where = ""
if condition is not None:
clause, filter_params = condition
if clause:
where = f" WHERE {clause}"
params.update(filter_params)

cls.cur.execute(
cls.query, (np.array(query.vector), top), binary=True, prepare=True
cls.query_template.format(where=where), params, binary=True, prepare=True
)
return cls.cur.fetchall()

Expand Down
12 changes: 7 additions & 5 deletions engine/clients/pgvector/upload.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import numpy as np
import psycopg
from pgvector.psycopg import register_vector
from psycopg.types.json import Jsonb

from dataset_reader.base_reader import Record
from engine.base_client import IncompatibilityError
Expand All @@ -29,19 +30,20 @@ def init_client(cls, host, distance, connection_params, upload_params):

@classmethod
def upload_batch(cls, batch: List[Record]):
ids, vectors = [], []
ids, vectors, payloads = [], [], []
for record in batch:
ids.append(record.id)
vectors.append(record.vector)
payloads.append(Jsonb(record.metadata))

vectors = np.array(vectors)
# Copy is faster than insert
with cls.cur.copy(
"COPY items (id, embedding) FROM STDIN WITH (FORMAT BINARY)"
"COPY items (id, embedding, payload) FROM STDIN WITH (FORMAT BINARY)"
) as copy:
copy.set_types(["integer", "vector"])
for i, embedding in zip(ids, vectors):
copy.write_row((i, embedding))
copy.set_types(["integer", "vector", "jsonb"])
for i, embedding, payload in zip(ids, vectors, payloads):
copy.write_row((i, embedding, payload))

@classmethod
def post_upload(cls, distance):
Expand Down
78 changes: 78 additions & 0 deletions tests/engine/clients/pgvector/test_pgvector_parser.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
import pytest
from psycopg.types.json import Jsonb

from engine.base_client import IncompatibilityError
from engine.clients.pgvector.parser import PgVectorConditionParser


@pytest.fixture
def pgvector_condition_parser():
return PgVectorConditionParser()


def test_parse_returns_none_on_none(pgvector_condition_parser):
assert pgvector_condition_parser.parse(None) is None


def test_parse_returns_none_on_empty(pgvector_condition_parser):
assert pgvector_condition_parser.parse({}) is None


def test_parse_converts_exact_match(pgvector_condition_parser):
conditions = {"and": [{"product_group_name": {"match": {"value": "Shoes"}}}]}
clause, params = pgvector_condition_parser.parse(conditions)

# `@>` so a list-valued payload field ("contains") matches too, not just scalars
assert "( payload -> %(p1)s @> %(p2)s )" == clause
assert "product_group_name" == params["p1"]
assert isinstance(params["p2"], Jsonb)
assert "Shoes" == params["p2"].obj


def test_parse_converts_multiple_or_statements(pgvector_condition_parser):
conditions = {
"or": [{"a": {"match": {"value": 80}}}, {"a": {"match": {"value": 2}}}]
}
clause, params = pgvector_condition_parser.parse(conditions)

assert "( payload -> %(p1)s @> %(p2)s OR payload -> %(p3)s @> %(p4)s )" == clause
assert "a" == params["p1"] == params["p3"]
assert 80 == params["p2"].obj
assert 2 == params["p4"].obj


def test_parse_converts_range(pgvector_condition_parser):
conditions = {"and": [{"price": {"range": {"gte": 10, "lt": 20}}}]}
clause, params = pgvector_condition_parser.parse(conditions)

# bounds are emitted in (lt, gt, lte, gte) order
assert (
"( ( payload -> %(p1)s < %(p2)s AND payload -> %(p3)s >= %(p4)s ) )" == clause
)
assert "price" == params["p1"] == params["p3"]
assert 20 == params["p2"].obj
assert 10 == params["p4"].obj


def test_parse_combines_and_and_or(pgvector_condition_parser):
conditions = {
"and": [{"a": {"match": {"value": 1}}}],
"or": [{"b": {"match": {"value": 2}}}, {"b": {"match": {"value": 3}}}],
}
clause, params = pgvector_condition_parser.parse(conditions)

assert (
"( payload -> %(p3)s @> %(p4)s OR payload -> %(p5)s @> %(p6)s )"
" AND ( payload -> %(p1)s @> %(p2)s )" == clause
)
assert {"p1": "a", "p3": "b", "p5": "b"} == {
k: v for k, v in params.items() if not isinstance(v, Jsonb)
}


def test_parse_geo_raises_incompatibility(pgvector_condition_parser):
conditions = {
"and": [{"a": {"geo": {"lon": 116.0, "lat": -52.0, "radius": 326341}}}]
}
with pytest.raises(IncompatibilityError):
pgvector_condition_parser.parse(conditions)