feat(vector): add direct pgvector adapter
This commit is contained in:
@@ -0,0 +1,50 @@
|
||||
# Local pgvector Task 1 report
|
||||
|
||||
## Status
|
||||
|
||||
Implemented the direct `PgVectorStore` behind the transport-neutral `VectorStore` port.
|
||||
The adapter uses separate optional reader and writer database configurations, derives
|
||||
capabilities from configured authority, validates strict positive search limits, filters kinds
|
||||
in SQL before limiting, and merges multi-collection results by cosine similarity.
|
||||
|
||||
All collection identifiers are selected from the fixed `schema_records`, `evidence`, and
|
||||
`memory` allowlist and composed with `psycopg2.sql.Identifier`. Values, vectors, kinds, hashes,
|
||||
and limits remain bound parameters. Collection/kind mismatches fail with `VectorStoreError`.
|
||||
|
||||
Upserts preserve the canonical metadata shape, use `record_key` conflict semantics, update the
|
||||
transport hash and embedding, and leave semantic metadata fields intact. Health probes reader
|
||||
and writer independently and reports observed `vector(N)` dimensions against the configured
|
||||
embedding dimension.
|
||||
|
||||
## Configuration and factory
|
||||
|
||||
`pgvector_direct` now accepts explicit optional `reader` and `writer` `DatabaseConfig` entries.
|
||||
The former `connection` entry remains supported as a deprecated read-only compatibility path.
|
||||
`build_vector_store(..., require_write=True)` accepts writer-only direct configurations and
|
||||
fails early when no explicit writer is present.
|
||||
|
||||
The transitional `build_vector_loader` bulk-sync path remains in place. It uses an explicit
|
||||
direct writer when present, or the legacy `connection`; it deliberately does not treat a new
|
||||
reader-only credential as writable. No production schema migration was added.
|
||||
|
||||
## TDD and verification
|
||||
|
||||
- RED: the new tests initially failed at collection because `PgVectorStore` did not exist.
|
||||
- Docker L0 pgvector tests: `11 passed`.
|
||||
- Direct + HTTP parity/factory/config focus: `51 passed`.
|
||||
- Full harness: `461 passed, 5 deselected`.
|
||||
- Changed-file Ruff lint: clean.
|
||||
- Changed-file Ruff format check: clean.
|
||||
- `git diff --check`: clean.
|
||||
|
||||
The repository-wide `ruff check .` still reports 34 pre-existing test-file findings outside
|
||||
Task 1; none are in changed files. The full pytest suite emits 17 existing legacy-config
|
||||
deprecation warnings.
|
||||
|
||||
## Scope and concerns
|
||||
|
||||
- Test fixtures create only the three existing vector tables needed to exercise the adapter;
|
||||
migration/versioning remains Task 2.
|
||||
- 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
|
||||
transition.
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Central construction of deployment-specific adapters."""
|
||||
|
||||
from tht.adapters.dwh import PostgresDwhAdapter, ThothRestDwhAdapter
|
||||
from tht.adapters.vector import LegacyDirectVectorStore, ThothHttpVectorStore
|
||||
from tht.adapters.vector import PgVectorStore, ThothHttpVectorStore
|
||||
from tht.config import Config, ConfigError
|
||||
from tht.db.connection import make_engine
|
||||
from tht.ports.dwh import DwhAdapter
|
||||
@@ -32,13 +32,13 @@ def build_vector_store(cfg: Config, *, require_write: bool = False) -> VectorSto
|
||||
|
||||
match resource.type:
|
||||
case "pgvector_direct":
|
||||
if require_write:
|
||||
reader = resource.reader or resource.connection
|
||||
if require_write and resource.writer is None:
|
||||
raise ConfigError("Vector writer non configurato per pgvector_direct")
|
||||
dim = cfg.embeddings.dim if cfg.embeddings is not None else 768
|
||||
return LegacyDirectVectorStore(
|
||||
make_engine(resource.connection),
|
||||
schema=resource.connection.db_schema,
|
||||
dim=dim,
|
||||
return PgVectorStore(
|
||||
reader,
|
||||
resource.writer,
|
||||
expected_dimension=cfg.embeddings.dim if cfg.embeddings is not None else None,
|
||||
)
|
||||
case "thoth_vector_http":
|
||||
if require_write and resource.writer is None:
|
||||
@@ -60,15 +60,19 @@ def build_vector_loader(cfg: Config, collection: str):
|
||||
if cfg.embeddings is None:
|
||||
raise ConfigError("Embeddings non configurati")
|
||||
|
||||
if resource.type == "thoth_vector_http" and resource.writer is not None and (
|
||||
cfg.profile == "workstation" or resource.direct is None
|
||||
if (
|
||||
resource.type == "thoth_vector_http"
|
||||
and resource.writer is not None
|
||||
and (cfg.profile == "workstation" or resource.direct is None)
|
||||
):
|
||||
from tht.vectorstore.rest_writer import RestVectorWriter
|
||||
|
||||
return RestVectorWriter(VectorRestClient(resource.writer), table=collection)
|
||||
|
||||
connection = (
|
||||
resource.connection if resource.type == "pgvector_direct" else resource.direct
|
||||
resource.writer or resource.connection
|
||||
if resource.type == "pgvector_direct"
|
||||
else resource.direct
|
||||
)
|
||||
if connection is None:
|
||||
raise ConfigError("Vector writer non configurato")
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""Vector-store adapter implementations."""
|
||||
|
||||
from tht.adapters.vector.legacy_direct import LegacyDirectVectorStore
|
||||
from tht.adapters.vector.pgvector import PgVectorStore
|
||||
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
|
||||
|
||||
__all__ = ["LegacyDirectVectorStore", "ThothHttpVectorStore"]
|
||||
__all__ = ["LegacyDirectVectorStore", "PgVectorStore", "ThothHttpVectorStore"]
|
||||
|
||||
@@ -0,0 +1,235 @@
|
||||
"""Direct PostgreSQL/pgvector implementation of the vector port."""
|
||||
|
||||
import json
|
||||
import re
|
||||
|
||||
from psycopg2 import sql
|
||||
from sqlalchemy import Engine
|
||||
|
||||
from tht.config import DatabaseConfig
|
||||
from tht.db.connection import make_engine
|
||||
from tht.ports.vector import (
|
||||
VectorCapabilities,
|
||||
VectorHealth,
|
||||
VectorReadUnavailable,
|
||||
VectorStoreError,
|
||||
VectorWriteRecord,
|
||||
VectorWriteUnavailable,
|
||||
require_positive_limit,
|
||||
)
|
||||
from tht.vectorstore.store import VectorHit, hit_from_metadata
|
||||
|
||||
|
||||
COLLECTION_KINDS = {
|
||||
"schema_records": {"schema_table", "schema_column"},
|
||||
"evidence": {"evidence"},
|
||||
"memory": {"memory", "solved_question"},
|
||||
}
|
||||
ALLOWED_COLLECTIONS = frozenset(COLLECTION_KINDS)
|
||||
_VECTOR_DIMENSION = re.compile(r"^vector\((\d+)\)$")
|
||||
|
||||
|
||||
def _collection(schema: str, name: str) -> sql.Identifier:
|
||||
if name not in ALLOWED_COLLECTIONS:
|
||||
raise VectorStoreError(f"Collection not allowed: {name}")
|
||||
return sql.Identifier(schema, name)
|
||||
|
||||
|
||||
def _vector_literal(values: list[float]) -> str:
|
||||
return "[" + ",".join(str(float(value)) for value in values) + "]"
|
||||
|
||||
|
||||
def _validate_kinds(collection: str, kinds: list[str]) -> None:
|
||||
invalid = set(kinds) - COLLECTION_KINDS[collection]
|
||||
if invalid:
|
||||
raise VectorStoreError(f"Kind not allowed for {collection}: {', '.join(sorted(invalid))}")
|
||||
|
||||
|
||||
class PgVectorStore:
|
||||
"""Direct store with independent reader and writer database credentials."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
read_config: DatabaseConfig | None,
|
||||
write_config: DatabaseConfig | None = None,
|
||||
*,
|
||||
expected_dimension: int | None = None,
|
||||
):
|
||||
self._reader = make_engine(read_config) if read_config is not None else None
|
||||
self._writer = make_engine(write_config) if write_config is not None else None
|
||||
config = read_config or write_config
|
||||
self._schema = config.db_schema if config is not None else "vectors"
|
||||
if read_config and write_config and read_config.db_schema != write_config.db_schema:
|
||||
raise VectorStoreError("Reader and writer vector schemas must match")
|
||||
self._expected_dimension = expected_dimension
|
||||
|
||||
@property
|
||||
def capabilities(self) -> VectorCapabilities:
|
||||
writable = self._writer is not None
|
||||
return VectorCapabilities(
|
||||
search=self._reader is not None,
|
||||
existing_hashes=writable,
|
||||
upsert=writable,
|
||||
)
|
||||
|
||||
def _probe(self, engine: Engine | None) -> tuple[bool | None, str | None, set[int]]:
|
||||
if engine is None:
|
||||
return None, None, set()
|
||||
try:
|
||||
raw = engine.raw_connection()
|
||||
try:
|
||||
with raw.cursor() as cursor:
|
||||
cursor.execute("SELECT 1")
|
||||
cursor.execute(
|
||||
"""SELECT format_type(a.atttypid, a.atttypmod)
|
||||
FROM pg_attribute a
|
||||
JOIN pg_class c ON c.oid = a.attrelid
|
||||
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||
WHERE n.nspname = %s AND c.relname = ANY(%s)
|
||||
AND a.attname = 'embedding' AND NOT a.attisdropped""",
|
||||
(self._schema, list(ALLOWED_COLLECTIONS)),
|
||||
)
|
||||
dimensions = {
|
||||
int(match.group(1))
|
||||
for (type_name,) in cursor.fetchall()
|
||||
if (match := _VECTOR_DIMENSION.match(type_name))
|
||||
}
|
||||
return True, None, dimensions
|
||||
finally:
|
||||
raw.close()
|
||||
except Exception as exc:
|
||||
return False, str(exc), set()
|
||||
|
||||
def health(self) -> VectorHealth:
|
||||
read_ok, read_detail, read_dimensions = self._probe(self._reader)
|
||||
write_ok, write_detail, write_dimensions = self._probe(self._writer)
|
||||
dimensions = tuple(sorted(read_dimensions | write_dimensions))
|
||||
compatible = (
|
||||
None
|
||||
if self._expected_dimension is None or not dimensions
|
||||
else dimensions == (self._expected_dimension,)
|
||||
)
|
||||
reachable = [value for value in (read_ok, write_ok) if value is not None]
|
||||
details = [value for value in (read_detail, write_detail) if value]
|
||||
return VectorHealth(
|
||||
ok=bool(reachable) and all(reachable) and compatible is not False,
|
||||
detail="; ".join(details) or None,
|
||||
read_configured=self._reader is not None,
|
||||
read_reachable=read_ok,
|
||||
read_detail=read_detail,
|
||||
write_configured=self._writer is not None,
|
||||
write_reachable=write_ok,
|
||||
write_detail=write_detail,
|
||||
expected_dimension=self._expected_dimension,
|
||||
observed_dimensions=dimensions,
|
||||
dimension_compatible=compatible,
|
||||
)
|
||||
|
||||
def search(
|
||||
self,
|
||||
collections: list[str],
|
||||
embedding: list[float],
|
||||
*,
|
||||
limit: int,
|
||||
kinds: list[str] | None = None,
|
||||
) -> list[VectorHit]:
|
||||
require_positive_limit(limit)
|
||||
if self._reader is None:
|
||||
raise VectorReadUnavailable("Vector reader credential is not configured")
|
||||
if self._expected_dimension is not None and len(embedding) != self._expected_dimension:
|
||||
raise VectorStoreError("Query embedding dimension does not match configured dimension")
|
||||
hits: list[VectorHit] = []
|
||||
raw = self._reader.raw_connection()
|
||||
try:
|
||||
with raw.cursor() as cursor:
|
||||
for collection in collections:
|
||||
table = _collection(self._schema, collection)
|
||||
if kinds:
|
||||
_validate_kinds(collection, kinds)
|
||||
where = sql.SQL(" WHERE kind = ANY(%s)") if 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)
|
||||
params = [_vector_literal(embedding)]
|
||||
if kinds:
|
||||
params.append(kinds)
|
||||
params.extend((_vector_literal(embedding), limit))
|
||||
cursor.execute(query, params)
|
||||
hits.extend(hit_from_metadata(row[1], row[0]) for row in cursor.fetchall())
|
||||
finally:
|
||||
raw.close()
|
||||
return sorted(hits, key=lambda hit: hit.similarity, reverse=True)[:limit]
|
||||
|
||||
def _require_writer(self) -> Engine:
|
||||
if self._writer is None:
|
||||
raise VectorWriteUnavailable("Vector writer credential is not configured")
|
||||
return self._writer
|
||||
|
||||
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]:
|
||||
engine = self._require_writer()
|
||||
table = _collection(self._schema, collection)
|
||||
_validate_kinds(collection, kinds)
|
||||
raw = engine.raw_connection()
|
||||
try:
|
||||
with raw.cursor() as cursor:
|
||||
cursor.execute(
|
||||
sql.SQL("SELECT record_key, content_hash FROM {} WHERE kind = ANY(%s)").format(
|
||||
table
|
||||
),
|
||||
(kinds,),
|
||||
)
|
||||
return dict(cursor.fetchall())
|
||||
finally:
|
||||
raw.close()
|
||||
|
||||
def upsert(self, collection: str, records: list[VectorWriteRecord]) -> int:
|
||||
engine = self._require_writer()
|
||||
table = _collection(self._schema, collection)
|
||||
for write_record in records:
|
||||
_validate_kinds(collection, [write_record.record.kind])
|
||||
if (
|
||||
self._expected_dimension is not None
|
||||
and len(write_record.embedding) != self._expected_dimension
|
||||
):
|
||||
raise VectorStoreError("Embedding dimension does not match configured dimension")
|
||||
query = sql.SQL(
|
||||
"INSERT INTO {} (record_key, kind, content_hash, metadata, embedding) "
|
||||
"VALUES (%s, %s, %s, %s::jsonb, %s::vector) "
|
||||
"ON CONFLICT (record_key) DO UPDATE SET kind = EXCLUDED.kind, "
|
||||
"content_hash = EXCLUDED.content_hash, metadata = EXCLUDED.metadata, "
|
||||
"embedding = EXCLUDED.embedding, indexed_at = now()"
|
||||
).format(table)
|
||||
raw = engine.raw_connection()
|
||||
try:
|
||||
with raw.cursor() as cursor:
|
||||
for write_record in records:
|
||||
record = write_record.record
|
||||
metadata = {
|
||||
"kind": record.kind,
|
||||
"ref": record.ref,
|
||||
"record_key": record.id,
|
||||
"title": record.title,
|
||||
"content": record.content,
|
||||
**record.metadata,
|
||||
}
|
||||
cursor.execute(
|
||||
query,
|
||||
(
|
||||
record.id,
|
||||
record.kind,
|
||||
write_record.content_hash,
|
||||
json.dumps(metadata),
|
||||
_vector_literal(write_record.embedding),
|
||||
),
|
||||
)
|
||||
raw.commit()
|
||||
except Exception:
|
||||
raw.rollback()
|
||||
raise
|
||||
finally:
|
||||
raw.close()
|
||||
return len(records)
|
||||
|
||||
|
||||
__all__ = ["ALLOWED_COLLECTIONS", "PgVectorStore"]
|
||||
+30
-9
@@ -18,6 +18,7 @@ class ConfigError(Exception):
|
||||
|
||||
def _expand_env(value: Any) -> Any:
|
||||
if isinstance(value, str):
|
||||
|
||||
def repl(m: re.Match) -> str:
|
||||
var = m.group(1)
|
||||
if var not in os.environ:
|
||||
@@ -85,7 +86,16 @@ DwhResourceConfig = Annotated[
|
||||
|
||||
class PgvectorDirectConfig(BaseModel):
|
||||
type: Literal["pgvector_direct"]
|
||||
connection: DatabaseConfig
|
||||
reader: DatabaseConfig | None = None
|
||||
writer: DatabaseConfig | None = None
|
||||
# Deprecated compatibility: a single direct connection historically meant read-only.
|
||||
connection: DatabaseConfig | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_connections(self):
|
||||
if self.reader is None and self.writer is None and self.connection is None:
|
||||
raise ValueError("pgvector_direct requires a reader or writer connection")
|
||||
return self
|
||||
|
||||
|
||||
class ThothVectorHttpConfig(BaseModel):
|
||||
@@ -134,9 +144,9 @@ class LshConfig(BaseModel):
|
||||
|
||||
class EligibilityConfig(BaseModel):
|
||||
# Soglie del principio di column eligibility (testo ampio ignorato ovunque).
|
||||
max_declared_len: int = 128 # char/varchar dichiarati <= soglia: eligible senza campionare
|
||||
max_avg_length: int = 40 # fallback data-driven: lunghezza media valori campionati
|
||||
max_sampled_len: int = 200 # fallback data-driven: lunghezza massima valore campionato
|
||||
max_declared_len: int = 128 # char/varchar dichiarati <= soglia: eligible senza campionare
|
||||
max_avg_length: int = 40 # fallback data-driven: lunghezza media valori campionati
|
||||
max_sampled_len: int = 200 # fallback data-driven: lunghezza massima valore campionato
|
||||
# Colonne di servizio sempre ignorate per nome (match case-insensitive), a prescindere
|
||||
# dal tipo: metadati ETL/audit non analitici (es. timestamp di ultimo aggiornamento).
|
||||
ignore_columns: list[str] = ["etl_last_update"]
|
||||
@@ -168,7 +178,7 @@ class VectorConfig(BaseModel):
|
||||
|
||||
class SearchConfig(BaseModel):
|
||||
rrf_k: int = 60
|
||||
top_schema_tables: int = 12 # default `--top` per `tht search --kind schema` (n. tabelle)
|
||||
top_schema_tables: int = 12 # default `--top` per `tht search --kind schema` (n. tabelle)
|
||||
schema_chunk_pool: int = 150 # chunk tabella/colonna fusi prima dell'aggregazione a tabella
|
||||
|
||||
|
||||
@@ -181,9 +191,17 @@ class ExecutionConfig(BaseModel):
|
||||
max_aggregate_cells: int = 20
|
||||
max_export_rows: int = 100000
|
||||
forbidden_functions: list[str] = [
|
||||
"setval", "nextval", "pg_advisory_lock", "pg_advisory_xact_lock",
|
||||
"dblink", "dblink_exec", "pg_terminate_backend", "pg_cancel_backend",
|
||||
"lo_import", "lo_export", "pg_reload_conf",
|
||||
"setval",
|
||||
"nextval",
|
||||
"pg_advisory_lock",
|
||||
"pg_advisory_xact_lock",
|
||||
"dblink",
|
||||
"dblink_exec",
|
||||
"pg_terminate_backend",
|
||||
"pg_cancel_backend",
|
||||
"lo_import",
|
||||
"lo_export",
|
||||
"pg_reload_conf",
|
||||
]
|
||||
|
||||
|
||||
@@ -302,7 +320,10 @@ def _populate_legacy_views(raw: dict[str, Any]) -> None:
|
||||
vectors = raw.get("vectors")
|
||||
if isinstance(vectors, dict):
|
||||
if vectors.get("type") == "pgvector_direct":
|
||||
raw.setdefault("vector_db", vectors["connection"])
|
||||
raw.setdefault(
|
||||
"vector_db",
|
||||
vectors.get("writer") or vectors.get("reader") or vectors.get("connection"),
|
||||
)
|
||||
elif vectors.get("type") == "thoth_vector_http":
|
||||
raw.setdefault("vector_rest", vectors.get("reader"))
|
||||
raw.setdefault("vector_write_rest", vectors.get("writer"))
|
||||
|
||||
Reference in New Issue
Block a user