fix(vector): harden direct pgvector parity
This commit is contained in:
@@ -48,3 +48,28 @@ deprecation warnings.
|
|||||||
- The legacy single `connection` form stays read-only through the public port, matching its
|
- The legacy single `connection` form stays read-only through the public port, matching its
|
||||||
previous adapter behavior, while remaining available to the explicitly documented bulk-loader
|
previous adapter behavior, while remaining available to the explicitly documented bulk-loader
|
||||||
transition.
|
transition.
|
||||||
|
|
||||||
|
## Review fix wave
|
||||||
|
|
||||||
|
The Task 1 review findings were addressed in a follow-up TDD cycle:
|
||||||
|
|
||||||
|
- Search now validates requested kinds against the global known-kind set, intersects valid kinds
|
||||||
|
with each collection, and skips unrelated collections. A direct-versus-HTTP parity test covers
|
||||||
|
the multi-collection case.
|
||||||
|
- Health requires all three allowlisted tables, an `embedding vector(N)` column on every table,
|
||||||
|
the expected dimension on every table, and the appropriate read or write table privileges for
|
||||||
|
each configured side. Empty and partial schemas return deterministic, credential-free details;
|
||||||
|
unexpected database failures expose only their exception class.
|
||||||
|
- The Docker L0 fixture now provisions separate least-privilege reader and writer roles. Tests
|
||||||
|
prove the reader cannot insert, the writer cannot execute the cosine-search SELECT, and the
|
||||||
|
adapter still routes search to the reader and upsert/hash operations to the writer. Direct
|
||||||
|
upsert uses an atomic `INSERT ... ON CONFLICT DO NOTHING` followed by `UPDATE` for an existing
|
||||||
|
key, avoiding broad SELECT authority while retaining conflict-safe hash/upsert semantics.
|
||||||
|
|
||||||
|
Fresh verification after the fix wave:
|
||||||
|
|
||||||
|
- Docker L0 + HTTP port/search parity: `42 passed` (earlier checkpoint); the final L0 file has
|
||||||
|
`16 passed` including the stricter raw-role search denial.
|
||||||
|
- Expanded focused adapter/config suite: `56 passed`.
|
||||||
|
- Full harness: `466 passed, 5 deselected`.
|
||||||
|
- Changed-file Ruff lint/format and `git diff --check`: clean.
|
||||||
|
|||||||
@@ -1,7 +1,9 @@
|
|||||||
import pytest
|
import pytest
|
||||||
from sqlalchemy import create_engine
|
from sqlalchemy import create_engine, text
|
||||||
|
from sqlalchemy.exc import ProgrammingError
|
||||||
from testcontainers.postgres import PostgresContainer
|
from testcontainers.postgres import PostgresContainer
|
||||||
|
|
||||||
|
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
|
||||||
from tht.config import DatabaseConfig
|
from tht.config import DatabaseConfig
|
||||||
from tht.ports.vector import (
|
from tht.ports.vector import (
|
||||||
VectorReadUnavailable,
|
VectorReadUnavailable,
|
||||||
@@ -28,11 +30,11 @@ def _record(content_hash: str, embedding: list[float], *, kind: str = "memory"):
|
|||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="module")
|
@pytest.fixture(scope="module")
|
||||||
def vector_config():
|
def vector_configs():
|
||||||
with PostgresContainer("pgvector/pgvector:pg16") as pg:
|
with PostgresContainer("pgvector/pgvector:pg16") as pg:
|
||||||
host = pg.get_container_host_ip()
|
host = pg.get_container_host_ip()
|
||||||
port = int(pg.get_exposed_port(5432))
|
port = int(pg.get_exposed_port(5432))
|
||||||
config = DatabaseConfig(
|
admin_config = DatabaseConfig(
|
||||||
host=host,
|
host=host,
|
||||||
port=port,
|
port=port,
|
||||||
database=pg.dbname,
|
database=pg.dbname,
|
||||||
@@ -56,15 +58,41 @@ def vector_config():
|
|||||||
indexed_at timestamptz NOT NULL DEFAULT now()
|
indexed_at timestamptz NOT NULL DEFAULT now()
|
||||||
)
|
)
|
||||||
""")
|
""")
|
||||||
|
connection.exec_driver_sql("CREATE ROLE vector_l0_reader LOGIN PASSWORD 'reader'")
|
||||||
|
connection.exec_driver_sql("CREATE ROLE vector_l0_writer LOGIN PASSWORD 'writer'")
|
||||||
|
connection.exec_driver_sql(
|
||||||
|
"GRANT USAGE ON SCHEMA vectors TO vector_l0_reader, vector_l0_writer"
|
||||||
|
)
|
||||||
|
connection.exec_driver_sql(
|
||||||
|
"GRANT SELECT ON ALL TABLES IN SCHEMA vectors TO vector_l0_reader"
|
||||||
|
)
|
||||||
|
connection.exec_driver_sql(
|
||||||
|
"GRANT USAGE, SELECT ON ALL SEQUENCES IN SCHEMA vectors TO vector_l0_writer"
|
||||||
|
)
|
||||||
|
for table in ("schema_records", "evidence", "memory"):
|
||||||
|
connection.exec_driver_sql(
|
||||||
|
f"GRANT INSERT, UPDATE ON vectors.{table} TO vector_l0_writer"
|
||||||
|
)
|
||||||
|
connection.exec_driver_sql(
|
||||||
|
f"GRANT SELECT (record_key, kind, content_hash) "
|
||||||
|
f"ON vectors.{table} TO vector_l0_writer"
|
||||||
|
)
|
||||||
engine.dispose()
|
engine.dispose()
|
||||||
yield config
|
reader_config = admin_config.model_copy(
|
||||||
|
update={"user": "vector_l0_reader", "password": "reader"}
|
||||||
|
)
|
||||||
|
writer_config = admin_config.model_copy(
|
||||||
|
update={"user": "vector_l0_writer", "password": "writer"}
|
||||||
|
)
|
||||||
|
yield admin_config, reader_config, writer_config
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def store(vector_config):
|
def store(vector_configs):
|
||||||
from tht.adapters.vector.pgvector import PgVectorStore
|
from tht.adapters.vector.pgvector import PgVectorStore
|
||||||
|
|
||||||
store = PgVectorStore(vector_config, vector_config, expected_dimension=2)
|
_, reader_config, writer_config = vector_configs
|
||||||
|
store = PgVectorStore(reader_config, writer_config, expected_dimension=2)
|
||||||
store.upsert("memory", [_record("reset", [0.0, 1.0])])
|
store.upsert("memory", [_record("reset", [0.0, 1.0])])
|
||||||
yield store
|
yield store
|
||||||
|
|
||||||
@@ -88,6 +116,46 @@ def test_pgvector_search_filters_kinds_before_limit(store):
|
|||||||
assert hits[0].kind == "memory"
|
assert hits[0].kind == "memory"
|
||||||
|
|
||||||
|
|
||||||
|
def test_pgvector_multi_collection_search_skips_collections_unrelated_to_kinds(store):
|
||||||
|
store.upsert("evidence", [_record("evidence", [1.0, 0.0], kind="evidence")])
|
||||||
|
|
||||||
|
hits = store.search(["evidence", "memory"], [1.0, 0.0], limit=3, kinds=["memory"])
|
||||||
|
|
||||||
|
assert hits
|
||||||
|
assert {hit.kind for hit in hits} == {"memory"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_pgvector_multi_collection_kind_filter_matches_http_adapter(store):
|
||||||
|
class Reader:
|
||||||
|
def search_similar(self, collection, embedding, limit, kinds=None):
|
||||||
|
if collection != "memory" or "memory" not in (kinds or []):
|
||||||
|
return []
|
||||||
|
return [
|
||||||
|
{
|
||||||
|
"similarity": 1.0,
|
||||||
|
"metadata": {
|
||||||
|
"record_key": "record:a",
|
||||||
|
"kind": "memory",
|
||||||
|
"ref": "session:test",
|
||||||
|
"title": "a",
|
||||||
|
"content": "content a",
|
||||||
|
"content_hash": "a",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
direct = store.search(["evidence", "memory"], [1.0, 0.0], limit=1, kinds=["memory"])
|
||||||
|
http = ThothHttpVectorStore(Reader(), None).search(
|
||||||
|
["evidence", "memory"], [1.0, 0.0], limit=1, kinds=["memory"]
|
||||||
|
)
|
||||||
|
assert [(hit.id, hit.kind) for hit in direct] == [(hit.id, hit.kind) for hit in http]
|
||||||
|
|
||||||
|
|
||||||
|
def test_pgvector_search_rejects_unknown_kind_globally(store):
|
||||||
|
with pytest.raises(VectorStoreError, match="Kind not allowed"):
|
||||||
|
store.search(["memory"], [1.0, 0.0], limit=1, kinds=["unknown"])
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("limit", [True, False, 1.0, 0, -1])
|
@pytest.mark.parametrize("limit", [True, False, 1.0, 0, -1])
|
||||||
def test_pgvector_search_requires_strict_positive_limit(store, limit):
|
def test_pgvector_search_requires_strict_positive_limit(store, limit):
|
||||||
with pytest.raises(ValueError, match="positive integer"):
|
with pytest.raises(ValueError, match="positive integer"):
|
||||||
@@ -108,32 +176,108 @@ def test_pgvector_rejects_kinds_not_belonging_to_collection(store):
|
|||||||
store.upsert("evidence", [_record("wrong", [1.0, 0.0])])
|
store.upsert("evidence", [_record("wrong", [1.0, 0.0])])
|
||||||
|
|
||||||
|
|
||||||
def test_pgvector_separates_read_and_write_credentials(vector_config):
|
def test_pgvector_separates_read_and_write_credentials(vector_configs):
|
||||||
from tht.adapters.vector.pgvector import PgVectorStore
|
from tht.adapters.vector.pgvector import PgVectorStore
|
||||||
|
|
||||||
reader = PgVectorStore(vector_config, expected_dimension=2)
|
_, reader_config, writer_config = vector_configs
|
||||||
|
reader = PgVectorStore(reader_config, expected_dimension=2)
|
||||||
assert reader.capabilities.search is True
|
assert reader.capabilities.search is True
|
||||||
assert reader.capabilities.upsert is False
|
assert reader.capabilities.upsert is False
|
||||||
with pytest.raises(VectorWriteUnavailable):
|
with pytest.raises(VectorWriteUnavailable):
|
||||||
reader.upsert("memory", [])
|
reader.upsert("memory", [])
|
||||||
|
|
||||||
writer = PgVectorStore(None, vector_config, expected_dimension=2)
|
writer = PgVectorStore(None, writer_config, expected_dimension=2)
|
||||||
assert writer.capabilities.search is False
|
assert writer.capabilities.search is False
|
||||||
assert writer.capabilities.upsert is True
|
assert writer.capabilities.upsert is True
|
||||||
with pytest.raises(VectorReadUnavailable):
|
with pytest.raises(VectorReadUnavailable):
|
||||||
writer.search(["memory"], [1.0, 0.0], limit=1)
|
writer.search(["memory"], [1.0, 0.0], limit=1)
|
||||||
|
|
||||||
|
|
||||||
def test_pgvector_health_reports_dimension_and_each_connection(vector_config):
|
def test_pgvector_database_roles_are_least_privilege(vector_configs):
|
||||||
|
_, reader_config, writer_config = vector_configs
|
||||||
|
reader_engine = create_engine(
|
||||||
|
f"postgresql+psycopg2://{reader_config.user}:{reader_config.password}"
|
||||||
|
f"@{reader_config.host}:{reader_config.port}/{reader_config.database}"
|
||||||
|
)
|
||||||
|
writer_engine = create_engine(
|
||||||
|
f"postgresql+psycopg2://{writer_config.user}:{writer_config.password}"
|
||||||
|
f"@{writer_config.host}:{writer_config.port}/{writer_config.database}"
|
||||||
|
)
|
||||||
|
with pytest.raises(ProgrammingError):
|
||||||
|
with reader_engine.begin() as connection:
|
||||||
|
connection.execute(
|
||||||
|
text(
|
||||||
|
"INSERT INTO vectors.memory "
|
||||||
|
"(record_key, kind, content_hash, metadata, embedding) "
|
||||||
|
"VALUES ('forbidden', 'memory', 'x', '{}', '[1,0]')"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
with pytest.raises(ProgrammingError):
|
||||||
|
with writer_engine.connect() as connection:
|
||||||
|
connection.execute(
|
||||||
|
text(
|
||||||
|
"SELECT metadata, 1 - (embedding <=> '[1,0]'::vector) AS similarity "
|
||||||
|
"FROM vectors.memory ORDER BY embedding <=> '[1,0]'::vector LIMIT 1"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
reader_engine.dispose()
|
||||||
|
writer_engine.dispose()
|
||||||
|
|
||||||
|
|
||||||
|
def test_pgvector_health_reports_dimension_and_each_connection(vector_configs):
|
||||||
from tht.adapters.vector.pgvector import PgVectorStore
|
from tht.adapters.vector.pgvector import PgVectorStore
|
||||||
|
|
||||||
health = PgVectorStore(vector_config, vector_config, expected_dimension=2).health()
|
_, reader_config, writer_config = vector_configs
|
||||||
|
health = PgVectorStore(reader_config, writer_config, expected_dimension=2).health()
|
||||||
assert health.ok is True
|
assert health.ok is True
|
||||||
assert health.read_reachable is True
|
assert health.read_reachable is True
|
||||||
assert health.write_reachable is True
|
assert health.write_reachable is True
|
||||||
assert health.observed_dimensions == (2,)
|
assert health.observed_dimensions == (2,)
|
||||||
assert health.dimension_compatible is True
|
assert health.dimension_compatible is True
|
||||||
|
|
||||||
mismatch = PgVectorStore(vector_config, None, expected_dimension=3).health()
|
mismatch = PgVectorStore(reader_config, None, expected_dimension=3).health()
|
||||||
assert mismatch.ok is False
|
assert mismatch.ok is False
|
||||||
|
assert mismatch.read_reachable is False
|
||||||
|
assert mismatch.read_detail == (
|
||||||
|
"embedding dimension mismatch: evidence=2, memory=2, schema_records=2"
|
||||||
|
)
|
||||||
assert mismatch.dimension_compatible is False
|
assert mismatch.dimension_compatible is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_pgvector_health_rejects_clean_and_partial_schemas(vector_configs):
|
||||||
|
from tht.adapters.vector.pgvector import PgVectorStore
|
||||||
|
|
||||||
|
admin_config, _, _ = vector_configs
|
||||||
|
engine = create_engine(
|
||||||
|
f"postgresql+psycopg2://{admin_config.user}:{admin_config.password}"
|
||||||
|
f"@{admin_config.host}:{admin_config.port}/{admin_config.database}"
|
||||||
|
)
|
||||||
|
with engine.begin() as connection:
|
||||||
|
connection.exec_driver_sql("CREATE SCHEMA clean_vectors")
|
||||||
|
connection.exec_driver_sql("CREATE SCHEMA partial_vectors")
|
||||||
|
connection.exec_driver_sql(
|
||||||
|
"CREATE TABLE partial_vectors.memory "
|
||||||
|
"(record_key text, kind text, content_hash text, metadata jsonb)"
|
||||||
|
)
|
||||||
|
engine.dispose()
|
||||||
|
|
||||||
|
clean = PgVectorStore(
|
||||||
|
admin_config.model_copy(update={"db_schema": "clean_vectors"}),
|
||||||
|
expected_dimension=2,
|
||||||
|
).health()
|
||||||
|
assert clean.ok is False
|
||||||
|
assert clean.read_reachable is False
|
||||||
|
assert clean.read_detail == (
|
||||||
|
"vector schema incomplete: missing tables evidence, memory, schema_records"
|
||||||
|
)
|
||||||
|
|
||||||
|
partial = PgVectorStore(
|
||||||
|
admin_config.model_copy(update={"db_schema": "partial_vectors"}),
|
||||||
|
expected_dimension=2,
|
||||||
|
).health()
|
||||||
|
assert partial.ok is False
|
||||||
|
assert partial.read_reachable is False
|
||||||
|
assert partial.read_detail == (
|
||||||
|
"vector schema incomplete: missing tables evidence, schema_records; "
|
||||||
|
"missing embedding columns memory"
|
||||||
|
)
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ COLLECTION_KINDS = {
|
|||||||
"memory": {"memory", "solved_question"},
|
"memory": {"memory", "solved_question"},
|
||||||
}
|
}
|
||||||
ALLOWED_COLLECTIONS = frozenset(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"^vector\((\d+)\)$")
|
||||||
|
|
||||||
|
|
||||||
@@ -39,12 +40,18 @@ def _vector_literal(values: list[float]) -> str:
|
|||||||
return "[" + ",".join(str(float(value)) for value in values) + "]"
|
return "[" + ",".join(str(float(value)) for value in values) + "]"
|
||||||
|
|
||||||
|
|
||||||
def _validate_kinds(collection: str, kinds: list[str]) -> None:
|
def _validate_collection_kinds(collection: str, kinds: list[str]) -> None:
|
||||||
invalid = set(kinds) - COLLECTION_KINDS[collection]
|
invalid = set(kinds) - COLLECTION_KINDS[collection]
|
||||||
if invalid:
|
if invalid:
|
||||||
raise VectorStoreError(f"Kind not allowed for {collection}: {', '.join(sorted(invalid))}")
|
raise VectorStoreError(f"Kind not allowed for {collection}: {', '.join(sorted(invalid))}")
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_known_kinds(kinds: list[str]) -> None:
|
||||||
|
invalid = set(kinds) - ALLOWED_KINDS
|
||||||
|
if invalid:
|
||||||
|
raise VectorStoreError(f"Kind not allowed: {', '.join(sorted(invalid))}")
|
||||||
|
|
||||||
|
|
||||||
class PgVectorStore:
|
class PgVectorStore:
|
||||||
"""Direct store with independent reader and writer database credentials."""
|
"""Direct store with independent reader and writer database credentials."""
|
||||||
|
|
||||||
@@ -72,7 +79,9 @@ class PgVectorStore:
|
|||||||
upsert=writable,
|
upsert=writable,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _probe(self, engine: Engine | None) -> tuple[bool | None, str | None, set[int]]:
|
def _probe(
|
||||||
|
self, engine: Engine | None, *, writable: bool
|
||||||
|
) -> tuple[bool | None, str | None, set[int]]:
|
||||||
if engine is None:
|
if engine is None:
|
||||||
return None, None, set()
|
return None, None, set()
|
||||||
try:
|
try:
|
||||||
@@ -81,28 +90,86 @@ class PgVectorStore:
|
|||||||
with raw.cursor() as cursor:
|
with raw.cursor() as cursor:
|
||||||
cursor.execute("SELECT 1")
|
cursor.execute("SELECT 1")
|
||||||
cursor.execute(
|
cursor.execute(
|
||||||
"""SELECT format_type(a.atttypid, a.atttypmod)
|
"""SELECT c.relname, format_type(a.atttypid, a.atttypmod),
|
||||||
FROM pg_attribute a
|
has_table_privilege(current_user, c.oid, 'SELECT'),
|
||||||
JOIN pg_class c ON c.oid = a.attrelid
|
has_table_privilege(current_user, c.oid, 'INSERT'),
|
||||||
|
has_table_privilege(current_user, c.oid, 'UPDATE'),
|
||||||
|
has_column_privilege(current_user, c.oid, 'record_key', 'SELECT')
|
||||||
|
AND has_column_privilege(
|
||||||
|
current_user, c.oid, 'content_hash', 'SELECT'
|
||||||
|
)
|
||||||
|
AND has_column_privilege(current_user, c.oid, 'kind', 'SELECT')
|
||||||
|
FROM pg_class c
|
||||||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||||
|
LEFT JOIN pg_attribute a ON a.attrelid = c.oid
|
||||||
|
AND a.attname = 'embedding' AND NOT a.attisdropped
|
||||||
WHERE n.nspname = %s AND c.relname = ANY(%s)
|
WHERE n.nspname = %s AND c.relname = ANY(%s)
|
||||||
AND a.attname = 'embedding' AND NOT a.attisdropped""",
|
AND c.relkind IN ('r', 'p')""",
|
||||||
(self._schema, list(ALLOWED_COLLECTIONS)),
|
(self._schema, list(ALLOWED_COLLECTIONS)),
|
||||||
)
|
)
|
||||||
|
rows = cursor.fetchall()
|
||||||
|
present = {row[0] for row in rows}
|
||||||
|
missing_tables = sorted(ALLOWED_COLLECTIONS - present)
|
||||||
|
missing_embeddings = sorted(row[0] for row in rows if row[1] is None)
|
||||||
|
privilege_missing = sorted(
|
||||||
|
row[0]
|
||||||
|
for row in rows
|
||||||
|
if (writable and not (row[3] and row[4] and row[5]))
|
||||||
|
or (not writable and not row[2])
|
||||||
|
)
|
||||||
|
problems = []
|
||||||
|
if missing_tables:
|
||||||
|
problems.append("missing tables " + ", ".join(missing_tables))
|
||||||
|
if missing_embeddings:
|
||||||
|
problems.append(
|
||||||
|
"missing embedding columns " + ", ".join(missing_embeddings)
|
||||||
|
)
|
||||||
|
if privilege_missing:
|
||||||
|
authority = "write" if writable else "read"
|
||||||
|
problems.append(
|
||||||
|
f"missing {authority} privileges " + ", ".join(privilege_missing)
|
||||||
|
)
|
||||||
|
if problems:
|
||||||
|
return False, "vector schema incomplete: " + "; ".join(problems), set()
|
||||||
dimensions = {
|
dimensions = {
|
||||||
int(match.group(1))
|
int(match.group(1))
|
||||||
for (type_name,) in cursor.fetchall()
|
for _, type_name, *_ in rows
|
||||||
if (match := _VECTOR_DIMENSION.match(type_name))
|
if (match := _VECTOR_DIMENSION.match(type_name))
|
||||||
}
|
}
|
||||||
|
invalid_types = sorted(
|
||||||
|
row[0]
|
||||||
|
for row in rows
|
||||||
|
if row[1] is not None and not _VECTOR_DIMENSION.match(row[1])
|
||||||
|
)
|
||||||
|
if invalid_types:
|
||||||
|
return (
|
||||||
|
False,
|
||||||
|
"vector schema incomplete: invalid embedding types "
|
||||||
|
+ ", ".join(invalid_types),
|
||||||
|
set(),
|
||||||
|
)
|
||||||
|
if self._expected_dimension is not None:
|
||||||
|
mismatches = sorted(
|
||||||
|
f"{name}={int(match.group(1))}"
|
||||||
|
for name, type_name, *_ in rows
|
||||||
|
if (match := _VECTOR_DIMENSION.match(type_name))
|
||||||
|
and int(match.group(1)) != self._expected_dimension
|
||||||
|
)
|
||||||
|
if mismatches:
|
||||||
|
return (
|
||||||
|
False,
|
||||||
|
"embedding dimension mismatch: " + ", ".join(mismatches),
|
||||||
|
dimensions,
|
||||||
|
)
|
||||||
return True, None, dimensions
|
return True, None, dimensions
|
||||||
finally:
|
finally:
|
||||||
raw.close()
|
raw.close()
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
return False, str(exc), set()
|
return False, f"vector database probe failed: {type(exc).__name__}", set()
|
||||||
|
|
||||||
def health(self) -> VectorHealth:
|
def health(self) -> VectorHealth:
|
||||||
read_ok, read_detail, read_dimensions = self._probe(self._reader)
|
read_ok, read_detail, read_dimensions = self._probe(self._reader, writable=False)
|
||||||
write_ok, write_detail, write_dimensions = self._probe(self._writer)
|
write_ok, write_detail, write_dimensions = self._probe(self._writer, writable=True)
|
||||||
dimensions = tuple(sorted(read_dimensions | write_dimensions))
|
dimensions = tuple(sorted(read_dimensions | write_dimensions))
|
||||||
compatible = (
|
compatible = (
|
||||||
None
|
None
|
||||||
@@ -138,22 +205,27 @@ class PgVectorStore:
|
|||||||
raise VectorReadUnavailable("Vector reader credential is not configured")
|
raise VectorReadUnavailable("Vector reader credential is not configured")
|
||||||
if self._expected_dimension is not None and len(embedding) != self._expected_dimension:
|
if self._expected_dimension is not None and len(embedding) != self._expected_dimension:
|
||||||
raise VectorStoreError("Query embedding dimension does not match configured dimension")
|
raise VectorStoreError("Query embedding dimension does not match configured dimension")
|
||||||
|
if kinds:
|
||||||
|
_validate_known_kinds(kinds)
|
||||||
hits: list[VectorHit] = []
|
hits: list[VectorHit] = []
|
||||||
raw = self._reader.raw_connection()
|
raw = self._reader.raw_connection()
|
||||||
try:
|
try:
|
||||||
with raw.cursor() as cursor:
|
with raw.cursor() as cursor:
|
||||||
for collection in collections:
|
for collection in collections:
|
||||||
table = _collection(self._schema, collection)
|
table = _collection(self._schema, collection)
|
||||||
if kinds:
|
collection_kinds = (
|
||||||
_validate_kinds(collection, kinds)
|
sorted(set(kinds) & COLLECTION_KINDS[collection]) if kinds else None
|
||||||
where = sql.SQL(" WHERE kind = ANY(%s)") if kinds else sql.SQL("")
|
)
|
||||||
|
if kinds and not collection_kinds:
|
||||||
|
continue
|
||||||
|
where = sql.SQL(" WHERE kind = ANY(%s)") if collection_kinds else sql.SQL("")
|
||||||
query = sql.SQL(
|
query = sql.SQL(
|
||||||
"SELECT metadata, 1 - (embedding <=> %s::vector) AS similarity "
|
"SELECT metadata, 1 - (embedding <=> %s::vector) AS similarity "
|
||||||
"FROM {}{} ORDER BY embedding <=> %s::vector LIMIT %s"
|
"FROM {}{} ORDER BY embedding <=> %s::vector LIMIT %s"
|
||||||
).format(table, where)
|
).format(table, where)
|
||||||
params = [_vector_literal(embedding)]
|
params = [_vector_literal(embedding)]
|
||||||
if kinds:
|
if collection_kinds:
|
||||||
params.append(kinds)
|
params.append(collection_kinds)
|
||||||
params.extend((_vector_literal(embedding), limit))
|
params.extend((_vector_literal(embedding), limit))
|
||||||
cursor.execute(query, params)
|
cursor.execute(query, params)
|
||||||
hits.extend(hit_from_metadata(row[1], row[0]) for row in cursor.fetchall())
|
hits.extend(hit_from_metadata(row[1], row[0]) for row in cursor.fetchall())
|
||||||
@@ -169,7 +241,7 @@ class PgVectorStore:
|
|||||||
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]:
|
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]:
|
||||||
engine = self._require_writer()
|
engine = self._require_writer()
|
||||||
table = _collection(self._schema, collection)
|
table = _collection(self._schema, collection)
|
||||||
_validate_kinds(collection, kinds)
|
_validate_collection_kinds(collection, kinds)
|
||||||
raw = engine.raw_connection()
|
raw = engine.raw_connection()
|
||||||
try:
|
try:
|
||||||
with raw.cursor() as cursor:
|
with raw.cursor() as cursor:
|
||||||
@@ -187,18 +259,20 @@ class PgVectorStore:
|
|||||||
engine = self._require_writer()
|
engine = self._require_writer()
|
||||||
table = _collection(self._schema, collection)
|
table = _collection(self._schema, collection)
|
||||||
for write_record in records:
|
for write_record in records:
|
||||||
_validate_kinds(collection, [write_record.record.kind])
|
_validate_collection_kinds(collection, [write_record.record.kind])
|
||||||
if (
|
if (
|
||||||
self._expected_dimension is not None
|
self._expected_dimension is not None
|
||||||
and len(write_record.embedding) != self._expected_dimension
|
and len(write_record.embedding) != self._expected_dimension
|
||||||
):
|
):
|
||||||
raise VectorStoreError("Embedding dimension does not match configured dimension")
|
raise VectorStoreError("Embedding dimension does not match configured dimension")
|
||||||
query = sql.SQL(
|
insert = sql.SQL(
|
||||||
"INSERT INTO {} (record_key, kind, content_hash, metadata, embedding) "
|
"INSERT INTO {} (record_key, kind, content_hash, metadata, embedding) "
|
||||||
"VALUES (%s, %s, %s, %s::jsonb, %s::vector) "
|
"VALUES (%s, %s, %s, %s::jsonb, %s::vector) "
|
||||||
"ON CONFLICT (record_key) DO UPDATE SET kind = EXCLUDED.kind, "
|
"ON CONFLICT (record_key) DO NOTHING"
|
||||||
"content_hash = EXCLUDED.content_hash, metadata = EXCLUDED.metadata, "
|
).format(table)
|
||||||
"embedding = EXCLUDED.embedding, indexed_at = now()"
|
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)
|
).format(table)
|
||||||
raw = engine.raw_connection()
|
raw = engine.raw_connection()
|
||||||
try:
|
try:
|
||||||
@@ -213,16 +287,23 @@ class PgVectorStore:
|
|||||||
"content": record.content,
|
"content": record.content,
|
||||||
**record.metadata,
|
**record.metadata,
|
||||||
}
|
}
|
||||||
|
metadata_json = json.dumps(metadata)
|
||||||
|
vector = _vector_literal(write_record.embedding)
|
||||||
cursor.execute(
|
cursor.execute(
|
||||||
query,
|
insert,
|
||||||
(
|
(record.id, record.kind, write_record.content_hash, metadata_json, vector),
|
||||||
record.id,
|
|
||||||
record.kind,
|
|
||||||
write_record.content_hash,
|
|
||||||
json.dumps(metadata),
|
|
||||||
_vector_literal(write_record.embedding),
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
|
if cursor.rowcount == 0:
|
||||||
|
cursor.execute(
|
||||||
|
update,
|
||||||
|
(
|
||||||
|
record.kind,
|
||||||
|
write_record.content_hash,
|
||||||
|
metadata_json,
|
||||||
|
vector,
|
||||||
|
record.id,
|
||||||
|
),
|
||||||
|
)
|
||||||
raw.commit()
|
raw.commit()
|
||||||
except Exception:
|
except Exception:
|
||||||
raw.rollback()
|
raw.rollback()
|
||||||
|
|||||||
Reference in New Issue
Block a user