feat(vector): add direct pgvector adapter

This commit is contained in:
2026-07-12 01:01:15 +02:00
parent ebdd3aa2c5
commit b09341f07e
7 changed files with 515 additions and 35 deletions
+139
View File
@@ -0,0 +1,139 @@
import pytest
from sqlalchemy import create_engine
from testcontainers.postgres import PostgresContainer
from tht.config import DatabaseConfig
from tht.ports.vector import (
VectorReadUnavailable,
VectorRecord,
VectorStoreError,
VectorWriteRecord,
VectorWriteUnavailable,
)
def _record(content_hash: str, embedding: list[float], *, kind: str = "memory"):
return VectorWriteRecord(
record=VectorRecord(
id=f"record:{content_hash}",
kind=kind,
ref="session:test",
title=content_hash,
content=f"content {content_hash}",
metadata={"content_hash": content_hash},
),
embedding=embedding,
content_hash=content_hash,
)
@pytest.fixture(scope="module")
def vector_config():
with PostgresContainer("pgvector/pgvector:pg16") as pg:
host = pg.get_container_host_ip()
port = int(pg.get_exposed_port(5432))
config = DatabaseConfig(
host=host,
port=port,
database=pg.dbname,
schema="vectors",
user=pg.username,
password=pg.password,
)
engine = create_engine(pg.get_connection_url())
with engine.begin() as connection:
connection.exec_driver_sql("CREATE EXTENSION vector")
connection.exec_driver_sql("CREATE SCHEMA vectors")
for table in ("schema_records", "evidence", "memory"):
connection.exec_driver_sql(f"""
CREATE TABLE vectors.{table} (
id bigserial PRIMARY KEY,
record_key text UNIQUE NOT NULL,
kind text NOT NULL,
content_hash text NOT NULL,
metadata jsonb NOT NULL,
embedding vector(2) NOT NULL,
indexed_at timestamptz NOT NULL DEFAULT now()
)
""")
engine.dispose()
yield config
@pytest.fixture
def store(vector_config):
from tht.adapters.vector.pgvector import PgVectorStore
store = PgVectorStore(vector_config, vector_config, expected_dimension=2)
store.upsert("memory", [_record("reset", [0.0, 1.0])])
yield store
def test_pgvector_round_trip_hash_and_upsert(store):
assert store.upsert("memory", [_record("a", [1.0, 0.0])]) == 1
assert store.existing_hashes("memory", ["memory"])["record:a"] == "a"
hits = store.search(["memory"], [1.0, 0.0], limit=5, kinds=["memory"])
assert hits[0].metadata["content_hash"] == "a"
assert hits[0].id == "record:a"
assert store.upsert("memory", [_record("a", [0.8, 0.2])]) == 1
assert store.search(["memory"], [0.8, 0.2], limit=1)[0].id == "record:a"
def test_pgvector_search_filters_kinds_before_limit(store):
store.upsert("memory", [_record("solved", [1.0, 0.0], kind="solved_question")])
hits = store.search("memory".split(), [1.0, 0.0], limit=1, kinds=["memory"])
assert len(hits) == 1
assert hits[0].kind == "memory"
@pytest.mark.parametrize("limit", [True, False, 1.0, 0, -1])
def test_pgvector_search_requires_strict_positive_limit(store, limit):
with pytest.raises(ValueError, match="positive integer"):
store.search(["memory"], [1.0, 0.0], limit=limit)
def test_pgvector_allowlists_collections(store):
with pytest.raises(VectorStoreError, match="Collection not allowed"):
store.search(["memory; DROP SCHEMA vectors"], [1.0, 0.0], limit=1)
with pytest.raises(VectorStoreError, match="Collection not allowed"):
store.upsert("unknown", [])
def test_pgvector_rejects_kinds_not_belonging_to_collection(store):
with pytest.raises(VectorStoreError, match="Kind not allowed"):
store.existing_hashes("evidence", ["memory"])
with pytest.raises(VectorStoreError, match="Kind not allowed"):
store.upsert("evidence", [_record("wrong", [1.0, 0.0])])
def test_pgvector_separates_read_and_write_credentials(vector_config):
from tht.adapters.vector.pgvector import PgVectorStore
reader = PgVectorStore(vector_config, expected_dimension=2)
assert reader.capabilities.search is True
assert reader.capabilities.upsert is False
with pytest.raises(VectorWriteUnavailable):
reader.upsert("memory", [])
writer = PgVectorStore(None, vector_config, expected_dimension=2)
assert writer.capabilities.search is False
assert writer.capabilities.upsert is True
with pytest.raises(VectorReadUnavailable):
writer.search(["memory"], [1.0, 0.0], limit=1)
def test_pgvector_health_reports_dimension_and_each_connection(vector_config):
from tht.adapters.vector.pgvector import PgVectorStore
health = PgVectorStore(vector_config, vector_config, expected_dimension=2).health()
assert health.ok is True
assert health.read_reachable is True
assert health.write_reachable is True
assert health.observed_dimensions == (2,)
assert health.dimension_compatible is True
mismatch = PgVectorStore(vector_config, None, expected_dimension=3).health()
assert mismatch.ok is False
assert mismatch.dimension_compatible is False
+45 -15
View File
@@ -1,7 +1,7 @@
import pytest
from tht.adapters.dwh import PostgresDwhAdapter, ThothRestDwhAdapter
from tht.adapters.vector import LegacyDirectVectorStore, ThothHttpVectorStore
from tht.adapters.vector import PgVectorStore, ThothHttpVectorStore
from tht.adapters.factory import build_dwh, build_vector_store
from tht.config import Config, ConfigError
@@ -28,7 +28,11 @@ def _config(*, dwh_type="thoth_rest", vector_type="thoth_vector_http", reader=Tr
vectors = (
{
"type": "thoth_vector_http",
**({"reader": {"base_url": "https://vectors.test/", "api_key": "reader"}} if reader else {}),
**(
{"reader": {"base_url": "https://vectors.test/", "api_key": "reader"}}
if reader
else {}
),
**(
{"writer": {"base_url": "https://vectors.test/", "api_key": "writer"}}
if writer
@@ -38,13 +42,32 @@ def _config(*, dwh_type="thoth_rest", vector_type="thoth_vector_http", reader=Tr
if vector_type == "thoth_vector_http"
else {
"type": "pgvector_direct",
"connection": {
"host": "vector-db",
"database": "postgres",
"schema": "vectors",
"user": "reader",
"password": "secret",
},
**(
{
"reader": {
"host": "vector-db",
"database": "postgres",
"schema": "vectors",
"user": "reader",
"password": "secret",
}
}
if reader
else {}
),
**(
{
"writer": {
"host": "vector-db",
"database": "postgres",
"schema": "vectors",
"user": "writer",
"password": "secret",
}
}
if writer
else {}
),
}
)
legacy_database = (
@@ -57,9 +80,7 @@ def _config(*, dwh_type="thoth_rest", vector_type="thoth_vector_http", reader=Tr
"transport": "rest",
}
)
return Config.model_validate(
{"dwh": dwh, "vectors": vectors, "database": legacy_database}
)
return Config.model_validate({"dwh": dwh, "vectors": vectors, "database": legacy_database})
@pytest.mark.parametrize(
@@ -87,14 +108,23 @@ def test_factory_builds_writer_only_http_vector_when_write_is_required():
assert store.capabilities.upsert is True
def test_factory_selects_direct_vector_reader():
config = _config(vector_type="pgvector_direct")
def test_factory_selects_direct_vector_store_and_requires_writer():
config = _config(vector_type="pgvector_direct", writer=False)
assert isinstance(build_vector_store(config), LegacyDirectVectorStore)
assert isinstance(build_vector_store(config), PgVectorStore)
with pytest.raises(ConfigError, match="writer"):
build_vector_store(config, require_write=True)
def test_factory_builds_writer_only_direct_vector_when_write_is_required():
store = build_vector_store(
_config(vector_type="pgvector_direct", reader=False), require_write=True
)
assert isinstance(store, PgVectorStore)
assert store.capabilities.search is False
assert store.capabilities.upsert is True
def test_factory_propagates_non_default_statement_timeout():
config = _config(dwh_type="postgres_direct")
config.execution.statement_timeout_ms = 12_345