diff --git a/README.md b/README.md index 9b2b5036..d8db8ff4 100644 --- a/README.md +++ b/README.md @@ -81,6 +81,9 @@ The following environment variables are required to run the application: - `PG_POOL_PRE_PING`: (Optional) Set to "False" to disable SQLAlchemy's pre-ping check. Default is "True". When enabled, the connection pool issues a lightweight `SELECT 1` before handing out a pooled connection, so stale connections dropped by a remote server or middlebox idle timeout are transparently replaced instead of surfacing as query errors. Recommended for any deployment that connects to a remote PostgreSQL instance (managed Postgres, connections that traverse a load balancer, etc.). - `PG_POOL_RECYCLE`: (Optional) Maximum age in seconds of a pooled connection before it is recycled. Default is "-1" (disabled). Set to a positive value when the server enforces a hard idle or max-lifetime limit (e.g. "1800" for a 30-minute cap). - `POSTGRES_SCHEMA`: (Optional) Prepend this schema to the Postgres `search_path` so langchain's pgvector tables live in (and are read from) it. Unset by default (uses the user's default schema, typically `public`). Useful when sharing a database with other services — create the schema out-of-band first (`CREATE SCHEMA IF NOT EXISTS ; GRANT USAGE, CREATE ON SCHEMA TO ;`); the RAG API will not create it for you and fails fast at startup if the schema is missing. `public` is always appended to the resulting search path so the `vector` data type stays resolvable when the extension was installed there (the common case). Multiple schemas may be supplied as a comma-separated list (e.g. `myapp,extensions`) when the `vector` extension lives in a non-`public` schema. +- `PGVECTOR_CREATE_LEGACY_INDEXES`: (Optional) Set to "True" to create the legacy `custom_id` and `cmetadata->>'file_id'` indexes on startup. Default is "False". +- `PGVECTOR_MIGRATE_CMETADATA_JSONB`: (Optional) Set to "True" to migrate `langchain_pg_embedding.cmetadata` from JSON to JSONB on startup. Default is "False". +- `PGVECTOR_CREATE_CMETADATA_GIN_INDEX`: (Optional) Set to "True" to create the `cmetadata` JSONB GIN index on startup. Default is "False". The index is created only when `cmetadata` is already JSONB; for a legacy JSON column, also enable `PGVECTOR_MIGRATE_CMETADATA_JSONB` or the index step is skipped. - `RAG_HOST`: (Optional) The hostname or IP address where the API server will run. Defaults to "0.0.0.0" - `RAG_PORT`: (Optional) The port number where the API server will run. Defaults to port 8000. - `JWT_SECRET`: (Optional) The secret key used for verifying JWT tokens for requests. diff --git a/app/services/database.py b/app/services/database.py index 32ba2a76..39e2e0f1 100644 --- a/app/services/database.py +++ b/app/services/database.py @@ -1,6 +1,16 @@ # app/services/database.py +import os + import asyncpg -from app.config import DSN, logger +from app.config import DSN, POSTGRES_SCHEMA, logger +from app.services.vector_store.factory import _build_search_path, _parse_schemas + + +def _env_flag(name: str, default: bool = False) -> bool: + value = os.getenv(name) + if value is None: + return default + return value.lower() in ("true", "1", "yes", "y", "on") class PSQLDatabase: @@ -9,7 +19,14 @@ class PSQLDatabase: @classmethod async def get_pool(cls): if cls.pool is None: - cls.pool = await asyncpg.create_pool(dsn=DSN) + pool_args = {"dsn": DSN} + if POSTGRES_SCHEMA: + schemas = _parse_schemas(POSTGRES_SCHEMA) + if schemas: + pool_args["server_settings"] = { + "search_path": _build_search_path(schemas) + } + cls.pool = await asyncpg.create_pool(**pool_args) return cls.pool @classmethod @@ -20,72 +37,100 @@ async def close_pool(cls): async def ensure_vector_indexes(): - """Ensure required indexes on langchain_pg_embedding and migrate cmetadata to JSONB. - - Runs at startup. Idempotent — safe to call repeatedly. - Operations: - 1. B-tree index on custom_id. - 2. Expression index on (cmetadata->>'file_id'). - 3. DDL migration: JSON -> JSONB for cmetadata (skipped if already JSONB). - 4. GIN index (jsonb_path_ops) on cmetadata for containment queries. - """ + """Ensure optional pgvector indexes/migrations when explicitly enabled.""" table_name = "langchain_pg_embedding" column_name = "custom_id" - # You might want to standardize the index naming convention index_name = f"idx_{table_name}_{column_name}" + create_legacy_indexes = _env_flag("PGVECTOR_CREATE_LEGACY_INDEXES") + migrate_cmetadata_jsonb = _env_flag("PGVECTOR_MIGRATE_CMETADATA_JSONB") + create_cmetadata_gin_index = _env_flag("PGVECTOR_CREATE_CMETADATA_GIN_INDEX") pool = await PSQLDatabase.get_pool() async with pool.acquire() as conn: - await conn.execute( - f""" - CREATE INDEX IF NOT EXISTS {index_name} ON {table_name} ({column_name}); - """ - ) - - # Expression index for (cmetadata->>'file_id') — critical for query - # performance. ExtendedPgVector overrides LangChain's default - # jsonb_path_match() to emit cmetadata->>'file_id' = ... which uses - # this B-tree index for fast equality lookups. - await conn.execute( - f""" - CREATE INDEX IF NOT EXISTS idx_{table_name}_file_id - ON {table_name} ((cmetadata->>'file_id')); - """ - ) - - # Migrate cmetadata from JSON to JSONB (idempotent — skipped if already JSONB). - # Rollback: ALTER TABLE langchain_pg_embedding ALTER COLUMN cmetadata TYPE JSON USING cmetadata::json; - # NOTE: table name is hardcoded below (not interpolated) to avoid SQL injection. - await conn.execute( - """ - DO $$ - BEGIN - IF EXISTS ( - SELECT 1 FROM information_schema.columns - WHERE table_name = 'langchain_pg_embedding' - AND table_schema = current_schema() - AND column_name = 'cmetadata' - AND data_type = 'json' - ) THEN - SET LOCAL lock_timeout = '10s'; - ALTER TABLE langchain_pg_embedding - ALTER COLUMN cmetadata TYPE JSONB USING cmetadata::jsonb; - END IF; - END - $$; + if create_legacy_indexes: + await conn.execute( + f""" + CREATE INDEX IF NOT EXISTS {index_name} ON {table_name} ({column_name}); """ - ) + ) - # GIN index on cmetadata for efficient JSONB filtering - await conn.execute( + # Expression index for (cmetadata->>'file_id') - critical for query + # performance. ExtendedPgVector emits cmetadata->>'file_id' = ... + # so this B-tree index supports fast equality lookups. + await conn.execute( + f""" + CREATE INDEX IF NOT EXISTS idx_{table_name}_file_id + ON {table_name} ((cmetadata->>'file_id')); """ - CREATE INDEX IF NOT EXISTS ix_cmetadata_gin - ON langchain_pg_embedding - USING gin (cmetadata jsonb_path_ops); - """ - ) + ) + else: + logger.info( + "Skipping legacy vector indexes; set PGVECTOR_CREATE_LEGACY_INDEXES=true to enable" + ) + + if migrate_cmetadata_jsonb: + await conn.execute( + """ + DO $$ + BEGIN + IF EXISTS ( + SELECT 1 FROM information_schema.columns + WHERE table_name = 'langchain_pg_embedding' + AND table_schema = current_schema() + AND column_name = 'cmetadata' + AND data_type = 'json' + ) THEN + SET LOCAL lock_timeout = '10s'; + ALTER TABLE langchain_pg_embedding + ALTER COLUMN cmetadata TYPE JSONB USING cmetadata::jsonb; + END IF; + END + $$; + """ + ) + else: + logger.info( + "Skipping cmetadata JSONB migration; set PGVECTOR_MIGRATE_CMETADATA_JSONB=true to enable" + ) + + if create_cmetadata_gin_index: + cmetadata_type = await conn.fetchval( + """ + SELECT data_type FROM information_schema.columns + WHERE table_name = 'langchain_pg_embedding' + AND table_schema = current_schema() + AND column_name = 'cmetadata'; + """ + ) + if cmetadata_type == "jsonb": + await conn.execute( + """ + CREATE INDEX IF NOT EXISTS ix_cmetadata_gin + ON langchain_pg_embedding + USING gin (cmetadata jsonb_path_ops); + """ + ) + elif cmetadata_type == "json": + logger.warning( + "Skipping cmetadata GIN index because cmetadata uses legacy JSON; " + "set PGVECTOR_MIGRATE_CMETADATA_JSONB=true to migrate it first" + ) + elif cmetadata_type is None: + logger.warning( + "Skipping cmetadata GIN index because langchain_pg_embedding.cmetadata " + "was not found in the current schema" + ) + else: + logger.warning( + "Skipping cmetadata GIN index because cmetadata has unsupported type %s", + cmetadata_type, + ) + else: + logger.info( + "Skipping cmetadata GIN index; set PGVECTOR_CREATE_CMETADATA_GIN_INDEX=true to enable" + ) - logger.info("Vector database indexes ensured") + logger.info("Vector database startup DDL checks complete") async def pg_health_check() -> bool: diff --git a/tests/integration/test_database.py b/tests/integration/test_database.py new file mode 100644 index 00000000..b3b7b68e --- /dev/null +++ b/tests/integration/test_database.py @@ -0,0 +1,77 @@ +import asyncpg +import pytest + +from app.services import database + + +@pytest.mark.integration +async def test_guarded_ddl_uses_schema_and_skips_json_until_migrated( + monkeypatch, pg_url, caplog +): + """Guarded DDL targets POSTGRES_SCHEMA and waits for a JSONB column.""" + schema = "test_guarded_startup_ddl" + dsn = pg_url.replace("postgresql+psycopg2://", "postgresql://") + admin_conn = await asyncpg.connect(dsn) + + try: + await admin_conn.execute(f"DROP SCHEMA IF EXISTS {schema} CASCADE") + await admin_conn.execute(f"CREATE SCHEMA {schema}") + await admin_conn.execute( + f""" + CREATE TABLE {schema}.langchain_pg_embedding ( + cmetadata JSON + ) + """ + ) + monkeypatch.setattr(database, "DSN", dsn) + monkeypatch.setattr(database, "POSTGRES_SCHEMA", schema) + monkeypatch.setattr(database.PSQLDatabase, "pool", None) + monkeypatch.delenv("PGVECTOR_CREATE_LEGACY_INDEXES", raising=False) + monkeypatch.delenv("PGVECTOR_MIGRATE_CMETADATA_JSONB", raising=False) + monkeypatch.setenv("PGVECTOR_CREATE_CMETADATA_GIN_INDEX", "true") + + await database.ensure_vector_indexes() + + pool = await database.PSQLDatabase.get_pool() + async with pool.acquire() as conn: + assert await conn.fetchval("SELECT current_schema()") == schema + index_exists = await conn.fetchval( + """ + SELECT EXISTS ( + SELECT 1 FROM pg_indexes + WHERE schemaname = current_schema() + AND indexname = 'ix_cmetadata_gin' + ) + """ + ) + assert index_exists is False + assert "uses legacy JSON" in caplog.text + assert "PGVECTOR_MIGRATE_CMETADATA_JSONB=true" in caplog.text + + monkeypatch.setenv("PGVECTOR_MIGRATE_CMETADATA_JSONB", "true") + await database.ensure_vector_indexes() + + async with pool.acquire() as conn: + column_type = await conn.fetchval( + """ + SELECT data_type FROM information_schema.columns + WHERE table_schema = current_schema() + AND table_name = 'langchain_pg_embedding' + AND column_name = 'cmetadata' + """ + ) + index_exists = await conn.fetchval( + """ + SELECT EXISTS ( + SELECT 1 FROM pg_indexes + WHERE schemaname = current_schema() + AND indexname = 'ix_cmetadata_gin' + ) + """ + ) + assert column_type == "jsonb" + assert index_exists is True + finally: + await database.PSQLDatabase.close_pool() + await admin_conn.execute(f"DROP SCHEMA IF EXISTS {schema} CASCADE") + await admin_conn.close() diff --git a/tests/services/test_database.py b/tests/services/test_database.py index 18570482..be9059ff 100644 --- a/tests/services/test_database.py +++ b/tests/services/test_database.py @@ -1,17 +1,22 @@ import asyncio +from unittest.mock import AsyncMock import pytest -from app.services.database import ensure_vector_indexes, PSQLDatabase +from app.services import database +from app.services.database import PSQLDatabase, ensure_vector_indexes class CapturingConnection: """Records every SQL statement passed to execute().""" - def __init__(self): + def __init__(self, cmetadata_type="jsonb"): self.statements = [] + self.queries = [] + self.cmetadata_type = cmetadata_type - async def fetchval(self, query, index_name): - return False + async def fetchval(self, query): + self.queries.append(query) + return self.cmetadata_type async def execute(self, query): self.statements.append(query) @@ -37,9 +42,37 @@ def acquire(self): return CapturingAcquire(self._conn) -def _run_with_captured_conn(monkeypatch): +DDL_FLAGS = ( + "PGVECTOR_CREATE_LEGACY_INDEXES", + "PGVECTOR_MIGRATE_CMETADATA_JSONB", + "PGVECTOR_CREATE_CMETADATA_GIN_INDEX", +) + + +def test_get_pool_uses_configured_schema_search_path(monkeypatch): + expected_pool = object() + create_pool = AsyncMock(return_value=expected_pool) + monkeypatch.setattr(database.asyncpg, "create_pool", create_pool) + monkeypatch.setattr(database, "POSTGRES_SCHEMA", "myapp, extensions") + monkeypatch.setattr(PSQLDatabase, "pool", None) + + pool = asyncio.run(PSQLDatabase.get_pool()) + + assert pool is expected_pool + create_pool.assert_awaited_once_with( + dsn=database.DSN, + server_settings={"search_path": "myapp,extensions,public"}, + ) + + +def _run_with_captured_conn(monkeypatch, *enabled_flags, cmetadata_type="jsonb"): """Run ensure_vector_indexes() and return the captured connection.""" - conn = CapturingConnection() + for flag in DDL_FLAGS: + monkeypatch.delenv(flag, raising=False) + for flag in enabled_flags: + monkeypatch.setenv(flag, "true") + + conn = CapturingConnection(cmetadata_type=cmetadata_type) pool = CapturingPool(conn) async def fake_get_pool(): @@ -52,19 +85,27 @@ async def fake_get_pool(): def test_ensure_vector_indexes(monkeypatch): conn = _run_with_captured_conn(monkeypatch) - assert len(conn.statements) > 0 + assert conn.statements == [] + + +def test_ensure_vector_indexes_legacy_indexes_opt_in(monkeypatch): + conn = _run_with_captured_conn(monkeypatch, "PGVECTOR_CREATE_LEGACY_INDEXES") + + assert len(conn.statements) == 2 + assert "custom_id" in conn.statements[0] + assert "cmetadata->>'file_id'" in conn.statements[1] def test_ensure_vector_indexes_do_block_dollar_quoting(monkeypatch): """DO block must use $$ dollar-quoting, not single $.""" - conn = _run_with_captured_conn(monkeypatch) + conn = _run_with_captured_conn(monkeypatch, "PGVECTOR_MIGRATE_CMETADATA_JSONB") do_block = next(s for s in conn.statements if "DO" in s) assert "$$" in do_block, "DO block must use $$ dollar-quoting" def test_ensure_vector_indexes_jsonb_migration_sql(monkeypatch): """Migration block contains the correct ALTER COLUMN and schema filter.""" - conn = _run_with_captured_conn(monkeypatch) + conn = _run_with_captured_conn(monkeypatch, "PGVECTOR_MIGRATE_CMETADATA_JSONB") do_block = next(s for s in conn.statements if "DO" in s) assert "TYPE JSONB" in do_block assert "cmetadata::jsonb" in do_block @@ -73,14 +114,38 @@ def test_ensure_vector_indexes_jsonb_migration_sql(monkeypatch): def test_ensure_vector_indexes_lock_timeout(monkeypatch): """Migration sets a lock_timeout before ALTER TABLE.""" - conn = _run_with_captured_conn(monkeypatch) + conn = _run_with_captured_conn(monkeypatch, "PGVECTOR_MIGRATE_CMETADATA_JSONB") do_block = next(s for s in conn.statements if "DO" in s) assert "lock_timeout" in do_block def test_ensure_vector_indexes_gin_index(monkeypatch): """GIN index with jsonb_path_ops is created.""" - conn = _run_with_captured_conn(monkeypatch) + conn = _run_with_captured_conn(monkeypatch, "PGVECTOR_CREATE_CMETADATA_GIN_INDEX") gin_stmt = next(s for s in conn.statements if "ix_cmetadata_gin" in s) assert "jsonb_path_ops" in gin_stmt assert "USING gin" in gin_stmt + assert any("data_type" in query for query in conn.queries) + + +def test_ensure_vector_indexes_gin_index_warns_for_legacy_json(monkeypatch, caplog): + conn = _run_with_captured_conn( + monkeypatch, + "PGVECTOR_CREATE_CMETADATA_GIN_INDEX", + cmetadata_type="json", + ) + + assert conn.statements == [] + assert "uses legacy JSON" in caplog.text + assert "PGVECTOR_MIGRATE_CMETADATA_JSONB=true" in caplog.text + + +def test_ensure_vector_indexes_gin_index_warns_for_missing_column(monkeypatch, caplog): + conn = _run_with_captured_conn( + monkeypatch, + "PGVECTOR_CREATE_CMETADATA_GIN_INDEX", + cmetadata_type=None, + ) + + assert conn.statements == [] + assert "was not found in the current schema" in caplog.text