Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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 <name>; GRANT USAGE, CREATE ON SCHEMA <name> TO <app_user>;`); 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.
Expand Down
161 changes: 103 additions & 58 deletions app/services/database.py
Original file line number Diff line number Diff line change
@@ -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:
Expand All @@ -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
Expand All @@ -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:
Expand Down
77 changes: 77 additions & 0 deletions tests/integration/test_database.py
Original file line number Diff line number Diff line change
@@ -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()
Loading
Loading