astrocyte-postgres 0.12.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- astrocyte_postgres/__init__.py +11 -0
- astrocyte_postgres/store.py +919 -0
- astrocyte_postgres/wiki_store.py +458 -0
- astrocyte_postgres-0.12.0.dist-info/METADATA +133 -0
- astrocyte_postgres-0.12.0.dist-info/RECORD +7 -0
- astrocyte_postgres-0.12.0.dist-info/WHEEL +4 -0
- astrocyte_postgres-0.12.0.dist-info/entry_points.txt +8 -0
|
@@ -0,0 +1,919 @@
|
|
|
1
|
+
"""VectorStore + DocumentStore backed by PostgreSQL with the pgvector extension.
|
|
2
|
+
|
|
3
|
+
``PostgresStore`` satisfies both the ``VectorStore`` *and* ``DocumentStore``
|
|
4
|
+
protocols. The same ``astrocyte_vectors`` table that stores embeddings also
|
|
5
|
+
carries a ``text_fts tsvector`` column (GIN-indexed, maintained by a trigger)
|
|
6
|
+
so that ``search_fulltext`` runs BM25-style ``ts_rank`` without a separate
|
|
7
|
+
Elasticsearch deployment.
|
|
8
|
+
|
|
9
|
+
This gives the recall pipeline the ``keyword`` strategy for free: when the
|
|
10
|
+
gateway or test code resolves ``document_store = pgvector``, ``parallel_retrieve``
|
|
11
|
+
can fuse lexical and semantic hits via RRF — exactly as Hindsight does with
|
|
12
|
+
its vector+lexical layer.
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
from __future__ import annotations
|
|
16
|
+
|
|
17
|
+
import asyncio
|
|
18
|
+
import json
|
|
19
|
+
import os
|
|
20
|
+
import re
|
|
21
|
+
from datetime import UTC, datetime
|
|
22
|
+
from typing import Any, ClassVar
|
|
23
|
+
|
|
24
|
+
import psycopg
|
|
25
|
+
from astrocyte.tenancy import fq_function, fq_table, get_current_schema
|
|
26
|
+
from astrocyte.types import (
|
|
27
|
+
Document,
|
|
28
|
+
DocumentFilters,
|
|
29
|
+
DocumentHit,
|
|
30
|
+
HealthStatus,
|
|
31
|
+
VectorFilters,
|
|
32
|
+
VectorHit,
|
|
33
|
+
VectorItem,
|
|
34
|
+
)
|
|
35
|
+
from pgvector.psycopg import register_vector_async
|
|
36
|
+
from psycopg.rows import dict_row
|
|
37
|
+
from psycopg.types.json import Json
|
|
38
|
+
from psycopg_pool import AsyncConnectionPool
|
|
39
|
+
|
|
40
|
+
_TABLE_SAFE = re.compile(r"^[a-zA-Z_][a-zA-Z0-9_]*$")
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def _sanitize_table(name: str) -> str:
|
|
44
|
+
if not _TABLE_SAFE.match(name):
|
|
45
|
+
raise ValueError(f"Invalid table name: {name!r}")
|
|
46
|
+
return name
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _split_metadata_list(value: object) -> list[str]:
|
|
50
|
+
if not isinstance(value, str) or not value:
|
|
51
|
+
return []
|
|
52
|
+
return [part for part in value.split("|") if part]
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
class PostgresStore:
|
|
56
|
+
"""Tier 1 vector store using `pgvector` cosine distance search."""
|
|
57
|
+
|
|
58
|
+
SPI_VERSION: ClassVar[int] = 1
|
|
59
|
+
|
|
60
|
+
def __init__(
|
|
61
|
+
self,
|
|
62
|
+
dsn: str | None = None,
|
|
63
|
+
table_name: str = "astrocyte_vectors",
|
|
64
|
+
embedding_dimensions: int = 128,
|
|
65
|
+
bootstrap_schema: bool = True,
|
|
66
|
+
**kwargs: Any,
|
|
67
|
+
) -> None:
|
|
68
|
+
self._dsn = dsn or os.environ.get("DATABASE_URL") or os.environ.get("ASTROCYTE_PG_DSN")
|
|
69
|
+
if not self._dsn:
|
|
70
|
+
raise ValueError(
|
|
71
|
+
"PostgresStore requires `dsn` in vector_store_config or DATABASE_URL / ASTROCYTE_PG_DSN",
|
|
72
|
+
)
|
|
73
|
+
self._table = _sanitize_table(table_name)
|
|
74
|
+
self._dim = int(embedding_dimensions)
|
|
75
|
+
if self._dim < 1:
|
|
76
|
+
raise ValueError("embedding_dimensions must be >= 1")
|
|
77
|
+
self._bootstrap_schema = bool(bootstrap_schema)
|
|
78
|
+
self._pool: AsyncConnectionPool | None = None
|
|
79
|
+
self._pool_lock = asyncio.Lock()
|
|
80
|
+
# Per-tenant-schema bootstrap tracking. A single PostgresStore instance
|
|
81
|
+
# can serve multiple tenants when the gateway sets ``_current_schema``
|
|
82
|
+
# per request, so the legacy single-shot ``_schema_ready`` flag has
|
|
83
|
+
# been replaced with a set of schemas that have had bootstrap applied.
|
|
84
|
+
# When ``bootstrap_schema=False`` (production), the set is pre-seeded
|
|
85
|
+
# with every schema the store has been touched from, effectively a
|
|
86
|
+
# no-op.
|
|
87
|
+
self._bootstrapped_schemas: set[str] = set()
|
|
88
|
+
self._schema_lock = asyncio.Lock()
|
|
89
|
+
|
|
90
|
+
def _fq(self, table: str | None = None) -> str:
|
|
91
|
+
"""Schema-qualify a table name using the current tenant context.
|
|
92
|
+
|
|
93
|
+
Defaults to ``self._table`` for the store's primary table; pass an
|
|
94
|
+
explicit name (e.g. ``"astrocyte_banks"``) for the cross-cutting
|
|
95
|
+
helper tables this store also writes to.
|
|
96
|
+
"""
|
|
97
|
+
return fq_table(table or self._table)
|
|
98
|
+
|
|
99
|
+
def _fq_func(self, function_name: str) -> str:
|
|
100
|
+
"""Schema-qualify a function/trigger-function name."""
|
|
101
|
+
return fq_function(function_name)
|
|
102
|
+
|
|
103
|
+
async def _ensure_pool(self) -> AsyncConnectionPool:
|
|
104
|
+
async with self._pool_lock:
|
|
105
|
+
if self._pool is None:
|
|
106
|
+
|
|
107
|
+
async def configure(conn: psycopg.AsyncConnection) -> None:
|
|
108
|
+
await conn.execute("SELECT 1")
|
|
109
|
+
# Pin search_path to ``public`` first so unqualified table
|
|
110
|
+
# writes/reads (``astrocyte_vectors``, ``astrocyte_banks``,
|
|
111
|
+
# etc.) always target the canonical migrated tables.
|
|
112
|
+
# Postgres defaults search_path to ``"$user", public`` which
|
|
113
|
+
# routes writes to ``<user>.<table>`` if the user-named
|
|
114
|
+
# schema exists — silently splitting data across schemas
|
|
115
|
+
# and breaking the entire benchmark when migrations only
|
|
116
|
+
# ran against ``public``.
|
|
117
|
+
await conn.execute('SET search_path = public, "$user"')
|
|
118
|
+
# register_vector_async needs the `vector` type. Skip until pgvector exists (quick path:
|
|
119
|
+
# /health can run before in-app DDL; runbook path: migrations already created the extension).
|
|
120
|
+
async with conn.cursor() as cur:
|
|
121
|
+
await cur.execute("SELECT 1 FROM pg_extension WHERE extname = 'vector'")
|
|
122
|
+
ext_present = await cur.fetchone()
|
|
123
|
+
if ext_present:
|
|
124
|
+
await register_vector_async(conn)
|
|
125
|
+
await conn.commit()
|
|
126
|
+
|
|
127
|
+
self._pool = AsyncConnectionPool(
|
|
128
|
+
conninfo=self._dsn,
|
|
129
|
+
configure=configure,
|
|
130
|
+
open=False,
|
|
131
|
+
min_size=2,
|
|
132
|
+
# Sized for parallel retain + concurrent PgQueuer workers
|
|
133
|
+
# (e.g. persona-compile tasks each call ``list_vectors``,
|
|
134
|
+
# ``store_vectors``, etc.). With 10 retain records in
|
|
135
|
+
# flight and PgQueuer running unbounded persona-compile
|
|
136
|
+
# jobs in the background, the previous max_size=10
|
|
137
|
+
# exhausted the pool and triggered cascading
|
|
138
|
+
# ``PoolTimeout`` errors. 40 leaves ~3 connections of
|
|
139
|
+
# headroom per concurrent unit.
|
|
140
|
+
max_size=40,
|
|
141
|
+
kwargs={"connect_timeout": 10},
|
|
142
|
+
)
|
|
143
|
+
await self._pool.open()
|
|
144
|
+
return self._pool
|
|
145
|
+
|
|
146
|
+
async def _ensure_schema(self, pool: AsyncConnectionPool) -> None:
|
|
147
|
+
"""Apply the dev/test schema for this store's table_name in the active tenant schema.
|
|
148
|
+
|
|
149
|
+
Only runs when ``bootstrap_schema=True`` (the default for tests using
|
|
150
|
+
per-test ``table_name`` strings that migrations cannot pre-create).
|
|
151
|
+
Production sets ``bootstrap_schema=False`` and relies entirely on
|
|
152
|
+
``migrations/`` applied by ``scripts/migrate.sh`` at deploy time.
|
|
153
|
+
|
|
154
|
+
**Per-tenant aware.** A single ``PostgresStore`` instance can serve
|
|
155
|
+
multiple tenants when the gateway sets ``_current_schema`` per
|
|
156
|
+
request; this method tracks which schemas have been bootstrapped and
|
|
157
|
+
runs DDL once per (schema, table_name) pair.
|
|
158
|
+
|
|
159
|
+
**Invariant: this method MUST produce the same schema as the SQL
|
|
160
|
+
migrations for ``table_name='astrocyte_vectors'``.** Each DDL block
|
|
161
|
+
below is annotated with the migration file it mirrors. If you add
|
|
162
|
+
DDL here, add the matching migration. If you change the migration,
|
|
163
|
+
update this method.
|
|
164
|
+
|
|
165
|
+
``register_vector_async`` is intentionally NOT called here — the
|
|
166
|
+
pool's per-connection ``configure`` callback handles vector-type
|
|
167
|
+
registration uniformly across both bootstrap modes.
|
|
168
|
+
"""
|
|
169
|
+
if not self._bootstrap_schema:
|
|
170
|
+
return
|
|
171
|
+
active_schema = get_current_schema()
|
|
172
|
+
# Cheap fast-path before grabbing the lock.
|
|
173
|
+
if active_schema in self._bootstrapped_schemas:
|
|
174
|
+
return
|
|
175
|
+
async with self._schema_lock:
|
|
176
|
+
if active_schema in self._bootstrapped_schemas:
|
|
177
|
+
return
|
|
178
|
+
# Resolve all qualified names ONCE up front so the giant DDL
|
|
179
|
+
# block below stays readable. Captures `get_current_schema()` at
|
|
180
|
+
# this moment so the whole bootstrap targets the same schema.
|
|
181
|
+
vectors = self._fq()
|
|
182
|
+
banks = self._fq("astrocyte_banks")
|
|
183
|
+
grants = self._fq("astrocyte_bank_access_grants")
|
|
184
|
+
temporal = self._fq("astrocyte_temporal_facts")
|
|
185
|
+
fts_func = self._fq_func(f"{self._table}_fts_update")
|
|
186
|
+
async with pool.connection() as conn:
|
|
187
|
+
# Schemas don't auto-create; create the target schema first
|
|
188
|
+
# if it doesn't exist (no-op for the default ``public``).
|
|
189
|
+
await conn.execute(f'CREATE SCHEMA IF NOT EXISTS "{active_schema}"')
|
|
190
|
+
# Mirrors 001_extension.sql. Extensions live in a single
|
|
191
|
+
# schema cluster-wide; CREATE IF NOT EXISTS is idempotent.
|
|
192
|
+
await conn.execute("CREATE EXTENSION IF NOT EXISTS vector")
|
|
193
|
+
# Mirrors 002_astrocytes_vectors.sql (with embedding_dimensions
|
|
194
|
+
# bound to ``self._dim`` instead of psql's :embedding_dimensions
|
|
195
|
+
# variable, plus the lifecycle/layer columns from
|
|
196
|
+
# 004_memory_layer.sql and 006_lifecycle_indexes.sql folded in
|
|
197
|
+
# so a fresh test table needs no follow-up ALTERs).
|
|
198
|
+
await conn.execute(
|
|
199
|
+
f"""
|
|
200
|
+
CREATE TABLE IF NOT EXISTS {vectors} (
|
|
201
|
+
id TEXT PRIMARY KEY,
|
|
202
|
+
bank_id TEXT NOT NULL,
|
|
203
|
+
embedding vector({self._dim}) NOT NULL,
|
|
204
|
+
text TEXT NOT NULL,
|
|
205
|
+
metadata JSONB,
|
|
206
|
+
tags TEXT[],
|
|
207
|
+
fact_type TEXT,
|
|
208
|
+
occurred_at TIMESTAMPTZ,
|
|
209
|
+
memory_layer TEXT,
|
|
210
|
+
retained_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
|
211
|
+
forgotten_at TIMESTAMPTZ
|
|
212
|
+
)
|
|
213
|
+
"""
|
|
214
|
+
)
|
|
215
|
+
# Mirrors 003_indexes.sql. Index names stay UNqualified —
|
|
216
|
+
# Postgres puts them in the same schema as the indexed table
|
|
217
|
+
# automatically and qualifying them here is a syntax error.
|
|
218
|
+
await conn.execute(
|
|
219
|
+
f"CREATE INDEX IF NOT EXISTS {self._table}_bank_idx ON {vectors} (bank_id)"
|
|
220
|
+
)
|
|
221
|
+
# Mirrors 006_lifecycle_indexes.sql.
|
|
222
|
+
await conn.execute(
|
|
223
|
+
f"""
|
|
224
|
+
CREATE INDEX IF NOT EXISTS {self._table}_bank_retained_idx
|
|
225
|
+
ON {vectors} (bank_id, retained_at DESC)
|
|
226
|
+
"""
|
|
227
|
+
)
|
|
228
|
+
await conn.execute(
|
|
229
|
+
f"""
|
|
230
|
+
CREATE INDEX IF NOT EXISTS {self._table}_bank_occurred_idx
|
|
231
|
+
ON {vectors} (bank_id, occurred_at DESC)
|
|
232
|
+
WHERE occurred_at IS NOT NULL
|
|
233
|
+
"""
|
|
234
|
+
)
|
|
235
|
+
await conn.execute(
|
|
236
|
+
f"""
|
|
237
|
+
CREATE INDEX IF NOT EXISTS {self._table}_bank_current_idx
|
|
238
|
+
ON {vectors} (bank_id)
|
|
239
|
+
WHERE forgotten_at IS NULL
|
|
240
|
+
"""
|
|
241
|
+
)
|
|
242
|
+
# Mirrors 010_hybrid_recall_indexes.sql. Index name aligned
|
|
243
|
+
# with the migration so bootstrap=True and bootstrap=False
|
|
244
|
+
# produce the same schema (was previously diverging:
|
|
245
|
+
# ``..._bank_fact_type_idx`` vs migration ``..._current_idx``).
|
|
246
|
+
await conn.execute(
|
|
247
|
+
f"""
|
|
248
|
+
CREATE INDEX IF NOT EXISTS {self._table}_bank_fact_type_current_idx
|
|
249
|
+
ON {vectors} (bank_id, fact_type)
|
|
250
|
+
WHERE forgotten_at IS NULL
|
|
251
|
+
"""
|
|
252
|
+
)
|
|
253
|
+
# Mirrors 005_banks_access.sql.
|
|
254
|
+
await conn.execute(
|
|
255
|
+
f"""
|
|
256
|
+
CREATE TABLE IF NOT EXISTS {banks} (
|
|
257
|
+
id TEXT PRIMARY KEY,
|
|
258
|
+
tenant_id TEXT,
|
|
259
|
+
display_name TEXT,
|
|
260
|
+
description TEXT,
|
|
261
|
+
metadata JSONB NOT NULL DEFAULT '{{}}'::jsonb,
|
|
262
|
+
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
|
263
|
+
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
|
264
|
+
archived_at TIMESTAMPTZ
|
|
265
|
+
)
|
|
266
|
+
"""
|
|
267
|
+
)
|
|
268
|
+
await conn.execute(
|
|
269
|
+
f"""
|
|
270
|
+
CREATE INDEX IF NOT EXISTS astrocyte_banks_tenant_idx
|
|
271
|
+
ON {banks} (tenant_id)
|
|
272
|
+
WHERE tenant_id IS NOT NULL
|
|
273
|
+
"""
|
|
274
|
+
)
|
|
275
|
+
await conn.execute(
|
|
276
|
+
f"""
|
|
277
|
+
CREATE TABLE IF NOT EXISTS {grants} (
|
|
278
|
+
id BIGSERIAL PRIMARY KEY,
|
|
279
|
+
bank_id TEXT NOT NULL,
|
|
280
|
+
principal TEXT NOT NULL,
|
|
281
|
+
permissions TEXT[] NOT NULL,
|
|
282
|
+
metadata JSONB NOT NULL DEFAULT '{{}}'::jsonb,
|
|
283
|
+
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
|
284
|
+
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
|
285
|
+
revoked_at TIMESTAMPTZ,
|
|
286
|
+
UNIQUE (bank_id, principal)
|
|
287
|
+
)
|
|
288
|
+
"""
|
|
289
|
+
)
|
|
290
|
+
await conn.execute(
|
|
291
|
+
f"""
|
|
292
|
+
CREATE INDEX IF NOT EXISTS astrocyte_bank_access_grants_principal_idx
|
|
293
|
+
ON {grants} (principal)
|
|
294
|
+
WHERE revoked_at IS NULL
|
|
295
|
+
"""
|
|
296
|
+
)
|
|
297
|
+
# 007_wiki_tables.sql is intentionally NOT mirrored here —
|
|
298
|
+
# those tables are owned by ``astrocyte_postgres.wiki_store``
|
|
299
|
+
# which has its own ``_ensure_schema()``.
|
|
300
|
+
#
|
|
301
|
+
# 009_entities_trigram_embedding.sql is intentionally NOT
|
|
302
|
+
# mirrored here — entity tables are owned by the
|
|
303
|
+
# entity-resolution module and other store adapters.
|
|
304
|
+
#
|
|
305
|
+
# Mirrors 008_entities_temporal.sql (temporal_facts table only;
|
|
306
|
+
# the entity_* tables in 008 are owned by the entity-resolution
|
|
307
|
+
# adapter, not this store).
|
|
308
|
+
await conn.execute(
|
|
309
|
+
f"""
|
|
310
|
+
CREATE TABLE IF NOT EXISTS {temporal} (
|
|
311
|
+
id BIGSERIAL PRIMARY KEY,
|
|
312
|
+
bank_id TEXT NOT NULL,
|
|
313
|
+
memory_id TEXT NOT NULL,
|
|
314
|
+
temporal_phrase TEXT NOT NULL,
|
|
315
|
+
anchor_time TIMESTAMPTZ,
|
|
316
|
+
resolved_start TIMESTAMPTZ,
|
|
317
|
+
resolved_end TIMESTAMPTZ,
|
|
318
|
+
resolved_date DATE,
|
|
319
|
+
date_granularity TEXT,
|
|
320
|
+
confidence DOUBLE PRECISION NOT NULL DEFAULT 1.0,
|
|
321
|
+
metadata JSONB NOT NULL DEFAULT '{{}}'::jsonb,
|
|
322
|
+
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
|
323
|
+
UNIQUE (bank_id, memory_id, temporal_phrase)
|
|
324
|
+
)
|
|
325
|
+
"""
|
|
326
|
+
)
|
|
327
|
+
# Mirrors 011_text_fts.sql (BM25/full-text column +
|
|
328
|
+
# GIN index + trigger function + trigger + backfill).
|
|
329
|
+
await conn.execute(
|
|
330
|
+
f"ALTER TABLE {vectors} ADD COLUMN IF NOT EXISTS text_fts tsvector"
|
|
331
|
+
)
|
|
332
|
+
await conn.execute(
|
|
333
|
+
f"""
|
|
334
|
+
CREATE INDEX IF NOT EXISTS {self._table}_fts_idx
|
|
335
|
+
ON {vectors} USING GIN (text_fts)
|
|
336
|
+
"""
|
|
337
|
+
)
|
|
338
|
+
await conn.execute(
|
|
339
|
+
f"""
|
|
340
|
+
CREATE OR REPLACE FUNCTION {fts_func}()
|
|
341
|
+
RETURNS trigger LANGUAGE plpgsql AS $$
|
|
342
|
+
BEGIN
|
|
343
|
+
NEW.text_fts := to_tsvector('english', COALESCE(NEW.text, ''));
|
|
344
|
+
RETURN NEW;
|
|
345
|
+
END;
|
|
346
|
+
$$
|
|
347
|
+
"""
|
|
348
|
+
)
|
|
349
|
+
# DROP-then-CREATE so the function body change above takes
|
|
350
|
+
# effect on existing tables (CREATE TRIGGER has no OR REPLACE).
|
|
351
|
+
await conn.execute(
|
|
352
|
+
f"DROP TRIGGER IF EXISTS {self._table}_fts_trigger ON {vectors}"
|
|
353
|
+
)
|
|
354
|
+
await conn.execute(
|
|
355
|
+
f"""
|
|
356
|
+
CREATE TRIGGER {self._table}_fts_trigger
|
|
357
|
+
BEFORE INSERT OR UPDATE OF text ON {vectors}
|
|
358
|
+
FOR EACH ROW EXECUTE FUNCTION {fts_func}()
|
|
359
|
+
"""
|
|
360
|
+
)
|
|
361
|
+
# Backfill existing rows that have NULL text_fts.
|
|
362
|
+
await conn.execute(
|
|
363
|
+
f"""
|
|
364
|
+
UPDATE {vectors}
|
|
365
|
+
SET text_fts = to_tsvector('english', COALESCE(text, ''))
|
|
366
|
+
WHERE text_fts IS NULL
|
|
367
|
+
"""
|
|
368
|
+
)
|
|
369
|
+
await conn.commit()
|
|
370
|
+
self._bootstrapped_schemas.add(active_schema)
|
|
371
|
+
|
|
372
|
+
async def store_vectors(self, items: list[VectorItem]) -> list[str]:
|
|
373
|
+
pool = await self._ensure_pool()
|
|
374
|
+
await self._ensure_schema(pool)
|
|
375
|
+
stored: list[str] = []
|
|
376
|
+
async with pool.connection() as conn:
|
|
377
|
+
async with conn.cursor() as cur:
|
|
378
|
+
for item in items:
|
|
379
|
+
if len(item.vector) != self._dim:
|
|
380
|
+
raise ValueError(
|
|
381
|
+
f"Vector length {len(item.vector)} != embedding_dimensions {self._dim}",
|
|
382
|
+
)
|
|
383
|
+
await self._upsert_bank(cur, item.bank_id)
|
|
384
|
+
await cur.execute(
|
|
385
|
+
f"""
|
|
386
|
+
INSERT INTO {self._fq()}
|
|
387
|
+
(
|
|
388
|
+
id, bank_id, embedding, text, metadata, tags, fact_type,
|
|
389
|
+
occurred_at, memory_layer, retained_at, forgotten_at
|
|
390
|
+
)
|
|
391
|
+
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, NULL)
|
|
392
|
+
ON CONFLICT (id) DO UPDATE SET
|
|
393
|
+
bank_id = EXCLUDED.bank_id,
|
|
394
|
+
embedding = EXCLUDED.embedding,
|
|
395
|
+
text = EXCLUDED.text,
|
|
396
|
+
metadata = EXCLUDED.metadata,
|
|
397
|
+
tags = EXCLUDED.tags,
|
|
398
|
+
fact_type = EXCLUDED.fact_type,
|
|
399
|
+
occurred_at = EXCLUDED.occurred_at,
|
|
400
|
+
memory_layer = EXCLUDED.memory_layer,
|
|
401
|
+
retained_at = EXCLUDED.retained_at,
|
|
402
|
+
forgotten_at = NULL
|
|
403
|
+
""",
|
|
404
|
+
(
|
|
405
|
+
item.id,
|
|
406
|
+
item.bank_id,
|
|
407
|
+
item.vector,
|
|
408
|
+
item.text,
|
|
409
|
+
Json(item.metadata) if item.metadata is not None else None,
|
|
410
|
+
item.tags,
|
|
411
|
+
item.fact_type,
|
|
412
|
+
item.occurred_at,
|
|
413
|
+
item.memory_layer,
|
|
414
|
+
item.retained_at or datetime.now(UTC),
|
|
415
|
+
),
|
|
416
|
+
)
|
|
417
|
+
await self._upsert_temporal_facts(cur, item)
|
|
418
|
+
stored.append(item.id)
|
|
419
|
+
return stored
|
|
420
|
+
|
|
421
|
+
async def _upsert_bank(self, cur: psycopg.AsyncCursor[Any], bank_id: str) -> None:
|
|
422
|
+
await cur.execute(
|
|
423
|
+
f"""
|
|
424
|
+
INSERT INTO {self._fq("astrocyte_banks")} (id, updated_at)
|
|
425
|
+
VALUES (%s, NOW())
|
|
426
|
+
ON CONFLICT (id) DO UPDATE SET updated_at = NOW()
|
|
427
|
+
""",
|
|
428
|
+
(bank_id,),
|
|
429
|
+
)
|
|
430
|
+
|
|
431
|
+
async def _upsert_temporal_facts(self, cur: psycopg.AsyncCursor[Any], item: VectorItem) -> None:
|
|
432
|
+
metadata = item.metadata or {}
|
|
433
|
+
phrases = _split_metadata_list(metadata.get("temporal_phrase"))
|
|
434
|
+
resolved_dates = _split_metadata_list(metadata.get("resolved_date"))
|
|
435
|
+
granularities = _split_metadata_list(metadata.get("date_granularity"))
|
|
436
|
+
if not phrases or not resolved_dates:
|
|
437
|
+
return
|
|
438
|
+
anchor = metadata.get("temporal_anchor")
|
|
439
|
+
for index, phrase in enumerate(phrases):
|
|
440
|
+
resolved = resolved_dates[index] if index < len(resolved_dates) else resolved_dates[0]
|
|
441
|
+
granularity = granularities[index] if index < len(granularities) else None
|
|
442
|
+
await cur.execute(
|
|
443
|
+
f"""
|
|
444
|
+
INSERT INTO {self._fq("astrocyte_temporal_facts")}
|
|
445
|
+
(bank_id, memory_id, temporal_phrase, anchor_time, resolved_date, date_granularity)
|
|
446
|
+
VALUES (%s, %s, %s, %s::timestamptz, %s::date, %s)
|
|
447
|
+
ON CONFLICT (bank_id, memory_id, temporal_phrase) DO UPDATE SET
|
|
448
|
+
anchor_time = EXCLUDED.anchor_time,
|
|
449
|
+
resolved_date = EXCLUDED.resolved_date,
|
|
450
|
+
date_granularity = EXCLUDED.date_granularity
|
|
451
|
+
""",
|
|
452
|
+
(item.bank_id, item.id, phrase, anchor, resolved, granularity),
|
|
453
|
+
)
|
|
454
|
+
|
|
455
|
+
async def search_similar(
|
|
456
|
+
self,
|
|
457
|
+
query_vector: list[float],
|
|
458
|
+
bank_id: str,
|
|
459
|
+
limit: int = 10,
|
|
460
|
+
filters: VectorFilters | None = None,
|
|
461
|
+
) -> list[VectorHit]:
|
|
462
|
+
if len(query_vector) != self._dim:
|
|
463
|
+
raise ValueError(
|
|
464
|
+
f"Query vector length {len(query_vector)} != embedding_dimensions {self._dim}",
|
|
465
|
+
)
|
|
466
|
+
pool = await self._ensure_pool()
|
|
467
|
+
await self._ensure_schema(pool)
|
|
468
|
+
|
|
469
|
+
where = ["bank_id = %s"]
|
|
470
|
+
params: list[Any] = [query_vector, bank_id]
|
|
471
|
+
if filters and filters.as_of:
|
|
472
|
+
where.append("retained_at <= %s")
|
|
473
|
+
where.append("(forgotten_at IS NULL OR forgotten_at > %s)")
|
|
474
|
+
params.extend([filters.as_of, filters.as_of])
|
|
475
|
+
else:
|
|
476
|
+
where.append("forgotten_at IS NULL")
|
|
477
|
+
if filters and filters.tags:
|
|
478
|
+
where.append("tags && %s::text[]")
|
|
479
|
+
params.append(filters.tags)
|
|
480
|
+
if filters and filters.fact_types:
|
|
481
|
+
where.append("fact_type = ANY(%s::text[])")
|
|
482
|
+
params.append(filters.fact_types)
|
|
483
|
+
params.extend([query_vector, limit])
|
|
484
|
+
|
|
485
|
+
where_sql = " AND ".join(where)
|
|
486
|
+
# Cosine distance `<=>`; map to a 0–1-ish score via (1 - distance).
|
|
487
|
+
sql = f"""
|
|
488
|
+
SELECT id, text, metadata, tags, fact_type, occurred_at, memory_layer, retained_at,
|
|
489
|
+
(1 - (embedding <=> %s::vector))::float AS score
|
|
490
|
+
FROM {self._fq()}
|
|
491
|
+
WHERE {where_sql}
|
|
492
|
+
ORDER BY embedding <=> %s::vector
|
|
493
|
+
LIMIT %s
|
|
494
|
+
"""
|
|
495
|
+
|
|
496
|
+
async with pool.connection() as conn:
|
|
497
|
+
async with conn.cursor(row_factory=dict_row) as cur:
|
|
498
|
+
await cur.execute(sql, params)
|
|
499
|
+
rows = await cur.fetchall()
|
|
500
|
+
|
|
501
|
+
hits: list[VectorHit] = []
|
|
502
|
+
for row in rows:
|
|
503
|
+
score = float(row["score"])
|
|
504
|
+
if score < 0.0:
|
|
505
|
+
score = 0.0
|
|
506
|
+
if score > 1.0:
|
|
507
|
+
score = 1.0
|
|
508
|
+
md = row["metadata"]
|
|
509
|
+
if isinstance(md, str):
|
|
510
|
+
md = json.loads(md)
|
|
511
|
+
hits.append(
|
|
512
|
+
VectorHit(
|
|
513
|
+
id=row["id"],
|
|
514
|
+
text=row["text"],
|
|
515
|
+
score=score,
|
|
516
|
+
metadata=md,
|
|
517
|
+
tags=list(row["tags"]) if row["tags"] else None,
|
|
518
|
+
fact_type=row["fact_type"],
|
|
519
|
+
occurred_at=row["occurred_at"],
|
|
520
|
+
memory_layer=row.get("memory_layer"),
|
|
521
|
+
retained_at=row.get("retained_at"),
|
|
522
|
+
)
|
|
523
|
+
)
|
|
524
|
+
return hits
|
|
525
|
+
|
|
526
|
+
async def search_hybrid_semantic_bm25(
|
|
527
|
+
self,
|
|
528
|
+
query_vector: list[float],
|
|
529
|
+
query: str,
|
|
530
|
+
bank_id: str,
|
|
531
|
+
limit: int = 10,
|
|
532
|
+
filters: VectorFilters | None = None,
|
|
533
|
+
) -> dict[str, list[VectorHit | DocumentHit]]:
|
|
534
|
+
"""Return semantic and BM25 hits with one SQL round trip.
|
|
535
|
+
|
|
536
|
+
This is the Postgres-specific Hindsight-style fast path. It preserves
|
|
537
|
+
the public VectorStore/DocumentStore methods as fallbacks while letting
|
|
538
|
+
the native pgvector stack avoid two separate pool checkouts on hot recall.
|
|
539
|
+
"""
|
|
540
|
+
if len(query_vector) != self._dim:
|
|
541
|
+
raise ValueError(
|
|
542
|
+
f"Query vector length {len(query_vector)} != embedding_dimensions {self._dim}",
|
|
543
|
+
)
|
|
544
|
+
if not query or not query.strip():
|
|
545
|
+
semantic = await self.search_similar(query_vector, bank_id, limit=limit, filters=filters)
|
|
546
|
+
return {"semantic": semantic, "keyword": []}
|
|
547
|
+
|
|
548
|
+
pool = await self._ensure_pool()
|
|
549
|
+
await self._ensure_schema(pool)
|
|
550
|
+
|
|
551
|
+
semantic_where = ["bank_id = %s"]
|
|
552
|
+
semantic_params: list[Any] = [bank_id]
|
|
553
|
+
keyword_where = ["bank_id = %s", "text_fts @@ plainto_tsquery('english', %s)"]
|
|
554
|
+
keyword_params: list[Any] = [bank_id, query]
|
|
555
|
+
|
|
556
|
+
if filters and filters.as_of:
|
|
557
|
+
semantic_where.append("retained_at <= %s")
|
|
558
|
+
semantic_where.append("(forgotten_at IS NULL OR forgotten_at > %s)")
|
|
559
|
+
semantic_params.extend([filters.as_of, filters.as_of])
|
|
560
|
+
keyword_where.append("retained_at <= %s")
|
|
561
|
+
keyword_where.append("(forgotten_at IS NULL OR forgotten_at > %s)")
|
|
562
|
+
keyword_params.extend([filters.as_of, filters.as_of])
|
|
563
|
+
else:
|
|
564
|
+
semantic_where.append("forgotten_at IS NULL")
|
|
565
|
+
keyword_where.append("forgotten_at IS NULL")
|
|
566
|
+
|
|
567
|
+
if filters and filters.tags:
|
|
568
|
+
semantic_where.append("tags && %s::text[]")
|
|
569
|
+
semantic_params.append(filters.tags)
|
|
570
|
+
keyword_where.append("tags && %s::text[]")
|
|
571
|
+
keyword_params.append(filters.tags)
|
|
572
|
+
|
|
573
|
+
if filters and filters.fact_types:
|
|
574
|
+
semantic_where.append("fact_type = ANY(%s::text[])")
|
|
575
|
+
semantic_params.append(filters.fact_types)
|
|
576
|
+
keyword_where.append("fact_type = ANY(%s::text[])")
|
|
577
|
+
keyword_params.append(filters.fact_types)
|
|
578
|
+
|
|
579
|
+
semantic_sql = " AND ".join(semantic_where)
|
|
580
|
+
keyword_sql = " AND ".join(keyword_where)
|
|
581
|
+
sql = f"""
|
|
582
|
+
WITH semantic AS (
|
|
583
|
+
SELECT
|
|
584
|
+
'semantic'::text AS strategy,
|
|
585
|
+
id, text, metadata, tags, fact_type, occurred_at,
|
|
586
|
+
memory_layer, retained_at,
|
|
587
|
+
(1 - (embedding <=> %s::vector))::float AS score
|
|
588
|
+
FROM {self._fq()}
|
|
589
|
+
WHERE {semantic_sql}
|
|
590
|
+
ORDER BY embedding <=> %s::vector
|
|
591
|
+
LIMIT %s
|
|
592
|
+
),
|
|
593
|
+
keyword AS (
|
|
594
|
+
SELECT
|
|
595
|
+
'keyword'::text AS strategy,
|
|
596
|
+
id, text, metadata, tags, fact_type, occurred_at,
|
|
597
|
+
memory_layer, retained_at,
|
|
598
|
+
ts_rank_cd(text_fts, plainto_tsquery('english', %s), 1)::float AS score
|
|
599
|
+
FROM {self._fq()}
|
|
600
|
+
WHERE {keyword_sql}
|
|
601
|
+
ORDER BY score DESC
|
|
602
|
+
LIMIT %s
|
|
603
|
+
)
|
|
604
|
+
SELECT * FROM semantic
|
|
605
|
+
UNION ALL
|
|
606
|
+
SELECT * FROM keyword
|
|
607
|
+
"""
|
|
608
|
+
params = [
|
|
609
|
+
query_vector,
|
|
610
|
+
*semantic_params,
|
|
611
|
+
query_vector,
|
|
612
|
+
limit,
|
|
613
|
+
query,
|
|
614
|
+
*keyword_params,
|
|
615
|
+
limit,
|
|
616
|
+
]
|
|
617
|
+
|
|
618
|
+
async with pool.connection() as conn:
|
|
619
|
+
async with conn.cursor(row_factory=dict_row) as cur:
|
|
620
|
+
await cur.execute(sql, params)
|
|
621
|
+
rows = await cur.fetchall()
|
|
622
|
+
|
|
623
|
+
semantic_hits: list[VectorHit] = []
|
|
624
|
+
keyword_hits: list[DocumentHit] = []
|
|
625
|
+
for row in rows:
|
|
626
|
+
md = row["metadata"]
|
|
627
|
+
if isinstance(md, str):
|
|
628
|
+
md = json.loads(md)
|
|
629
|
+
score = max(0.0, min(1.0, float(row["score"])))
|
|
630
|
+
if row["strategy"] == "semantic":
|
|
631
|
+
semantic_hits.append(
|
|
632
|
+
VectorHit(
|
|
633
|
+
id=row["id"],
|
|
634
|
+
text=row["text"],
|
|
635
|
+
score=score,
|
|
636
|
+
metadata=md,
|
|
637
|
+
tags=list(row["tags"]) if row["tags"] else None,
|
|
638
|
+
fact_type=row["fact_type"],
|
|
639
|
+
occurred_at=row["occurred_at"],
|
|
640
|
+
memory_layer=row.get("memory_layer"),
|
|
641
|
+
retained_at=row.get("retained_at"),
|
|
642
|
+
)
|
|
643
|
+
)
|
|
644
|
+
else:
|
|
645
|
+
keyword_hits.append(
|
|
646
|
+
DocumentHit(
|
|
647
|
+
document_id=row["id"],
|
|
648
|
+
text=row["text"],
|
|
649
|
+
score=score,
|
|
650
|
+
metadata=md,
|
|
651
|
+
)
|
|
652
|
+
)
|
|
653
|
+
|
|
654
|
+
return {"semantic": semantic_hits, "keyword": keyword_hits}
|
|
655
|
+
|
|
656
|
+
async def list_vectors(
|
|
657
|
+
self,
|
|
658
|
+
bank_id: str,
|
|
659
|
+
offset: int = 0,
|
|
660
|
+
limit: int = 100,
|
|
661
|
+
) -> list[VectorItem]:
|
|
662
|
+
pool = await self._ensure_pool()
|
|
663
|
+
await self._ensure_schema(pool)
|
|
664
|
+
async with pool.connection() as conn:
|
|
665
|
+
async with conn.cursor(row_factory=dict_row) as cur:
|
|
666
|
+
await cur.execute(
|
|
667
|
+
f"""
|
|
668
|
+
SELECT id, bank_id, embedding, text, metadata, tags, fact_type,
|
|
669
|
+
occurred_at, memory_layer, retained_at
|
|
670
|
+
FROM {self._fq()}
|
|
671
|
+
WHERE bank_id = %s
|
|
672
|
+
AND forgotten_at IS NULL
|
|
673
|
+
ORDER BY id
|
|
674
|
+
OFFSET %s LIMIT %s
|
|
675
|
+
""",
|
|
676
|
+
(bank_id, offset, limit),
|
|
677
|
+
)
|
|
678
|
+
rows = await cur.fetchall()
|
|
679
|
+
items: list[VectorItem] = []
|
|
680
|
+
for row in rows:
|
|
681
|
+
md = row["metadata"]
|
|
682
|
+
if isinstance(md, str):
|
|
683
|
+
md = json.loads(md)
|
|
684
|
+
items.append(
|
|
685
|
+
VectorItem(
|
|
686
|
+
id=row["id"],
|
|
687
|
+
bank_id=row["bank_id"],
|
|
688
|
+
vector=list(row["embedding"]),
|
|
689
|
+
text=row["text"],
|
|
690
|
+
metadata=md,
|
|
691
|
+
tags=list(row["tags"]) if row["tags"] else None,
|
|
692
|
+
fact_type=row["fact_type"],
|
|
693
|
+
occurred_at=row["occurred_at"],
|
|
694
|
+
memory_layer=row.get("memory_layer"),
|
|
695
|
+
retained_at=row.get("retained_at"),
|
|
696
|
+
)
|
|
697
|
+
)
|
|
698
|
+
return items
|
|
699
|
+
|
|
700
|
+
async def list_recent_vectors(
|
|
701
|
+
self,
|
|
702
|
+
bank_id: str,
|
|
703
|
+
limit: int = 100,
|
|
704
|
+
filters: VectorFilters | None = None,
|
|
705
|
+
) -> list[VectorItem]:
|
|
706
|
+
"""Return recent vectors using Postgres indexes instead of a Python scan."""
|
|
707
|
+
pool = await self._ensure_pool()
|
|
708
|
+
await self._ensure_schema(pool)
|
|
709
|
+
|
|
710
|
+
where = ["bank_id = %s"]
|
|
711
|
+
params: list[Any] = [bank_id]
|
|
712
|
+
if filters and filters.as_of:
|
|
713
|
+
where.append("retained_at <= %s")
|
|
714
|
+
where.append("(forgotten_at IS NULL OR forgotten_at > %s)")
|
|
715
|
+
params.extend([filters.as_of, filters.as_of])
|
|
716
|
+
else:
|
|
717
|
+
where.append("forgotten_at IS NULL")
|
|
718
|
+
if filters and filters.tags:
|
|
719
|
+
where.append("tags && %s::text[]")
|
|
720
|
+
params.append(filters.tags)
|
|
721
|
+
if filters and filters.fact_types:
|
|
722
|
+
where.append("fact_type = ANY(%s::text[])")
|
|
723
|
+
params.append(filters.fact_types)
|
|
724
|
+
if filters and filters.time_range:
|
|
725
|
+
start, end = filters.time_range
|
|
726
|
+
where.append("occurred_at >= %s")
|
|
727
|
+
where.append("occurred_at <= %s")
|
|
728
|
+
params.extend([start, end])
|
|
729
|
+
|
|
730
|
+
params.append(limit)
|
|
731
|
+
where_sql = " AND ".join(where)
|
|
732
|
+
async with pool.connection() as conn:
|
|
733
|
+
async with conn.cursor(row_factory=dict_row) as cur:
|
|
734
|
+
await cur.execute(
|
|
735
|
+
f"""
|
|
736
|
+
SELECT id, bank_id, embedding, text, metadata, tags, fact_type,
|
|
737
|
+
occurred_at, memory_layer, retained_at
|
|
738
|
+
FROM {self._fq()}
|
|
739
|
+
WHERE {where_sql}
|
|
740
|
+
ORDER BY COALESCE(occurred_at, retained_at) DESC, id
|
|
741
|
+
LIMIT %s
|
|
742
|
+
""",
|
|
743
|
+
params,
|
|
744
|
+
)
|
|
745
|
+
rows = await cur.fetchall()
|
|
746
|
+
|
|
747
|
+
items: list[VectorItem] = []
|
|
748
|
+
for row in rows:
|
|
749
|
+
md = row["metadata"]
|
|
750
|
+
if isinstance(md, str):
|
|
751
|
+
md = json.loads(md)
|
|
752
|
+
items.append(
|
|
753
|
+
VectorItem(
|
|
754
|
+
id=row["id"],
|
|
755
|
+
bank_id=row["bank_id"],
|
|
756
|
+
vector=list(row["embedding"]),
|
|
757
|
+
text=row["text"],
|
|
758
|
+
metadata=md,
|
|
759
|
+
tags=list(row["tags"]) if row["tags"] else None,
|
|
760
|
+
fact_type=row["fact_type"],
|
|
761
|
+
occurred_at=row["occurred_at"],
|
|
762
|
+
memory_layer=row.get("memory_layer"),
|
|
763
|
+
retained_at=row.get("retained_at"),
|
|
764
|
+
)
|
|
765
|
+
)
|
|
766
|
+
return items
|
|
767
|
+
|
|
768
|
+
async def delete(self, ids: list[str], bank_id: str) -> int:
|
|
769
|
+
if not ids:
|
|
770
|
+
return 0
|
|
771
|
+
pool = await self._ensure_pool()
|
|
772
|
+
await self._ensure_schema(pool)
|
|
773
|
+
async with pool.connection() as conn:
|
|
774
|
+
async with conn.cursor() as cur:
|
|
775
|
+
await cur.execute(
|
|
776
|
+
f"""
|
|
777
|
+
UPDATE {self._fq()}
|
|
778
|
+
SET forgotten_at = NOW()
|
|
779
|
+
WHERE bank_id = %s
|
|
780
|
+
AND id = ANY(%s::text[])
|
|
781
|
+
AND forgotten_at IS NULL
|
|
782
|
+
""",
|
|
783
|
+
(bank_id, ids),
|
|
784
|
+
)
|
|
785
|
+
return cur.rowcount or 0
|
|
786
|
+
|
|
787
|
+
async def close(self) -> None:
|
|
788
|
+
"""Close the connection pool. Safe to call multiple times."""
|
|
789
|
+
async with self._pool_lock:
|
|
790
|
+
if self._pool is not None:
|
|
791
|
+
await self._pool.close()
|
|
792
|
+
self._pool = None
|
|
793
|
+
|
|
794
|
+
async def __aenter__(self) -> "PostgresStore":
|
|
795
|
+
return self
|
|
796
|
+
|
|
797
|
+
async def __aexit__(self, *exc: object) -> None:
|
|
798
|
+
await self.close()
|
|
799
|
+
|
|
800
|
+
async def health(self) -> HealthStatus:
|
|
801
|
+
try:
|
|
802
|
+
pool = await self._ensure_pool()
|
|
803
|
+
async with pool.connection() as conn:
|
|
804
|
+
async with conn.cursor() as cur:
|
|
805
|
+
await cur.execute("SELECT 1")
|
|
806
|
+
return HealthStatus(healthy=True, message="pgvector connected")
|
|
807
|
+
except Exception as e:
|
|
808
|
+
return HealthStatus(healthy=False, message=f"pgvector unhealthy: {e!s}")
|
|
809
|
+
|
|
810
|
+
# ── DocumentStore protocol ────────────────────────────────────────────────
|
|
811
|
+
# PostgresStore satisfies DocumentStore so the recall pipeline can fuse
|
|
812
|
+
# lexical (BM25-style ts_rank) hits alongside semantic (cosine) hits via
|
|
813
|
+
# RRF — exactly the vector+lexical fusion Hindsight uses. The text is
|
|
814
|
+
# already stored in astrocyte_vectors at retain time; DocumentStore methods
|
|
815
|
+
# operate on the same table via the text_fts tsvector column.
|
|
816
|
+
|
|
817
|
+
async def store_document(self, document: Document, bank_id: str) -> str:
|
|
818
|
+
"""No-op: text is already stored by store_vectors() at retain time.
|
|
819
|
+
|
|
820
|
+
The tsvector trigger keeps text_fts in sync automatically. This
|
|
821
|
+
method exists to satisfy the DocumentStore protocol so that callers
|
|
822
|
+
(e.g. PipelineOrchestrator) can treat PostgresStore as a
|
|
823
|
+
DocumentStore without a separate code path.
|
|
824
|
+
"""
|
|
825
|
+
return document.id
|
|
826
|
+
|
|
827
|
+
async def search_fulltext(
|
|
828
|
+
self,
|
|
829
|
+
query: str,
|
|
830
|
+
bank_id: str,
|
|
831
|
+
limit: int = 10,
|
|
832
|
+
filters: DocumentFilters | None = None,
|
|
833
|
+
) -> list[DocumentHit]:
|
|
834
|
+
"""BM25-style full-text search using PostgreSQL ts_rank over text_fts.
|
|
835
|
+
|
|
836
|
+
Ranks results by ``ts_rank_cd`` (cover-density ranking), which
|
|
837
|
+
rewards query terms appearing close together — a good proxy for BM25
|
|
838
|
+
on the memory-text lengths typical of Astrocyte. Normalises the raw
|
|
839
|
+
score by document length (``|normalization|=1``) so long memories
|
|
840
|
+
don't dominate short ones.
|
|
841
|
+
"""
|
|
842
|
+
if not query or not query.strip():
|
|
843
|
+
return []
|
|
844
|
+
|
|
845
|
+
pool = await self._ensure_pool()
|
|
846
|
+
await self._ensure_schema(pool)
|
|
847
|
+
|
|
848
|
+
where = ["bank_id = %s", "forgotten_at IS NULL", "text_fts @@ plainto_tsquery('english', %s)"]
|
|
849
|
+
where_params: list[Any] = [bank_id, query]
|
|
850
|
+
|
|
851
|
+
if filters and filters.tags:
|
|
852
|
+
where.append("tags && %s::text[]")
|
|
853
|
+
where_params.append(filters.tags)
|
|
854
|
+
|
|
855
|
+
# SELECT ts_rank_cd(%s) appears before the WHERE %s bindings in the
|
|
856
|
+
# query string, so query must be the first positional param.
|
|
857
|
+
params = [query] + where_params + [limit]
|
|
858
|
+
|
|
859
|
+
where_sql = " AND ".join(where)
|
|
860
|
+
sql = f"""
|
|
861
|
+
SELECT
|
|
862
|
+
id,
|
|
863
|
+
text,
|
|
864
|
+
metadata,
|
|
865
|
+
ts_rank_cd(text_fts, plainto_tsquery('english', %s), 1) AS score
|
|
866
|
+
FROM {self._fq()}
|
|
867
|
+
WHERE {where_sql}
|
|
868
|
+
ORDER BY score DESC
|
|
869
|
+
LIMIT %s
|
|
870
|
+
"""
|
|
871
|
+
|
|
872
|
+
async with pool.connection() as conn:
|
|
873
|
+
async with conn.cursor(row_factory=dict_row) as cur:
|
|
874
|
+
await cur.execute(sql, params)
|
|
875
|
+
rows = await cur.fetchall()
|
|
876
|
+
|
|
877
|
+
hits: list[DocumentHit] = []
|
|
878
|
+
for row in rows:
|
|
879
|
+
md = row["metadata"]
|
|
880
|
+
if isinstance(md, str):
|
|
881
|
+
md = json.loads(md)
|
|
882
|
+
hits.append(
|
|
883
|
+
DocumentHit(
|
|
884
|
+
document_id=row["id"],
|
|
885
|
+
text=row["text"],
|
|
886
|
+
score=float(row["score"]),
|
|
887
|
+
metadata=md,
|
|
888
|
+
)
|
|
889
|
+
)
|
|
890
|
+
return hits
|
|
891
|
+
|
|
892
|
+
async def get_document(self, document_id: str, bank_id: str) -> Document | None:
|
|
893
|
+
"""Retrieve a stored memory as a Document by ID."""
|
|
894
|
+
pool = await self._ensure_pool()
|
|
895
|
+
await self._ensure_schema(pool)
|
|
896
|
+
|
|
897
|
+
async with pool.connection() as conn:
|
|
898
|
+
async with conn.cursor(row_factory=dict_row) as cur:
|
|
899
|
+
await cur.execute(
|
|
900
|
+
f"""
|
|
901
|
+
SELECT id, text, metadata, tags
|
|
902
|
+
FROM {self._fq()}
|
|
903
|
+
WHERE id = %s AND bank_id = %s AND forgotten_at IS NULL
|
|
904
|
+
""",
|
|
905
|
+
(document_id, bank_id),
|
|
906
|
+
)
|
|
907
|
+
row = await cur.fetchone()
|
|
908
|
+
|
|
909
|
+
if row is None:
|
|
910
|
+
return None
|
|
911
|
+
md = row["metadata"]
|
|
912
|
+
if isinstance(md, str):
|
|
913
|
+
md = json.loads(md)
|
|
914
|
+
return Document(
|
|
915
|
+
id=row["id"],
|
|
916
|
+
text=row["text"],
|
|
917
|
+
metadata=md,
|
|
918
|
+
tags=list(row["tags"]) if row["tags"] else None,
|
|
919
|
+
)
|