From 31d5312d9d7ada7c9611a92fd69dccf3f8b623d1 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 13 Sep 2026 17:13:13 +0000 Subject: [PATCH 01/10] feat(dialects): add pg_textsearch dialect and decouple paradedb --- pyproject.toml | 1 + sqlspec/core/splitter.py | 1 + sqlspec/dialects/__init__.py | 5 +- sqlspec/dialects/postgres/__init__.py | 3 +- sqlspec/dialects/postgres/_generators.py | 13 +++-- sqlspec/dialects/postgres/_operators.py | 11 ++++- sqlspec/dialects/postgres/_paradedb.py | 19 ++++++-- sqlspec/dialects/postgres/_pg_textsearch.py | 42 ++++++++++++++++ tests/unit/dialects/test_operators.py | 53 +++++++++++++++++++++ tests/unit/dialects/test_paradedb.py | 20 +++++--- tests/unit/dialects/test_pg_textsearch.py | 46 ++++++++++++++++++ tests/unit/dialects/test_registration.py | 5 +- 12 files changed, 200 insertions(+), 19 deletions(-) create mode 100644 sqlspec/dialects/postgres/_pg_textsearch.py create mode 100644 tests/unit/dialects/test_operators.py create mode 100644 tests/unit/dialects/test_pg_textsearch.py diff --git a/pyproject.toml b/pyproject.toml index 5ff276064..9ad2125de 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -161,6 +161,7 @@ sqlspec = "sqlspec.__main__:run_cli" [project.entry-points."sqlglot.dialects"] paradedb = "sqlspec.dialects.postgres:ParadeDB" +pg_textsearch = "sqlspec.dialects.postgres:PGTextSearch" pgvector = "sqlspec.dialects.postgres:PGVector" spangres = "sqlspec.dialects.spanner:Spangres" spanner = "sqlspec.dialects.spanner:Spanner" diff --git a/sqlspec/core/splitter.py b/sqlspec/core/splitter.py index e27868f6f..c84446435 100644 --- a/sqlspec/core/splitter.py +++ b/sqlspec/core/splitter.py @@ -551,6 +551,7 @@ class BigQueryDialectConfig(_EagerDialectConfig): "postgresql": PostgreSQLDialectConfig, "postgres": PostgreSQLDialectConfig, "paradedb": PostgreSQLDialectConfig, + "pg_textsearch": PostgreSQLDialectConfig, "pgvector": PostgreSQLDialectConfig, "mysql": MySQLDialectConfig, "sqlite": SQLiteDialectConfig, diff --git a/sqlspec/dialects/__init__.py b/sqlspec/dialects/__init__.py index aacab0fb1..e03dcbe7a 100644 --- a/sqlspec/dialects/__init__.py +++ b/sqlspec/dialects/__init__.py @@ -13,12 +13,13 @@ from typing import TYPE_CHECKING, Any if TYPE_CHECKING: - from sqlspec.dialects.postgres import ParadeDB, PGVector + from sqlspec.dialects.postgres import ParadeDB, PGTextSearch, PGVector from sqlspec.dialects.spanner import Spangres, Spanner -__all__ = ("PGVector", "ParadeDB", "Spangres", "Spanner") +__all__ = ("PGTextSearch", "PGVector", "ParadeDB", "Spangres", "Spanner") _DIALECT_MODULES = { + "PGTextSearch": "sqlspec.dialects.postgres", "PGVector": "sqlspec.dialects.postgres", "ParadeDB": "sqlspec.dialects.postgres", "Spangres": "sqlspec.dialects.spanner", diff --git a/sqlspec/dialects/postgres/__init__.py b/sqlspec/dialects/postgres/__init__.py index 503fb6a59..d66dff35f 100644 --- a/sqlspec/dialects/postgres/__init__.py +++ b/sqlspec/dialects/postgres/__init__.py @@ -4,6 +4,7 @@ """ from sqlspec.dialects.postgres._paradedb import ParadeDB +from sqlspec.dialects.postgres._pg_textsearch import PGTextSearch from sqlspec.dialects.postgres._pgvector import PGVector -__all__: tuple[str, ...] = ("PGVector", "ParadeDB") +__all__: tuple[str, ...] = ("PGTextSearch", "PGVector", "ParadeDB") diff --git a/sqlspec/dialects/postgres/_generators.py b/sqlspec/dialects/postgres/_generators.py index 73c286da1..e74c5dc79 100644 --- a/sqlspec/dialects/postgres/_generators.py +++ b/sqlspec/dialects/postgres/_generators.py @@ -7,6 +7,8 @@ interpreted subclasses. """ +from typing import Final + from sqlglot import exp from sqlglot.dialects.postgres import Postgres from sqlglot.generators.postgres import PostgresGenerator @@ -19,15 +21,19 @@ ) from sqlspec.dialects.postgres._operators import is_postgres_extension_operator, postgres_extension_operator -__all__ = ("PGVectorGenerator", "ParadeDBGenerator") +__all__ = ("PGTextSearchGenerator", "PGVectorGenerator", "ParadeDBGenerator") _BASE_OPERATOR_TRANSFORM = Postgres.Generator.TRANSFORMS[exp.Operator] +_POSTGRES_EXTENSION_DIALECT_NAMES: Final[frozenset[str]] = frozenset( + {"Postgres", "PGVector", "ParadeDB", "PGTextSearch"} +) def _postgres_extension_operator_sql(generator: PostgresGenerator, expression: exp.Operator) -> str: - dialect_class = getattr(generator.dialect, "__class__", None) + dialect = generator.dialect + dialect_class = getattr(dialect, "__class__", None) dialect_name = dialect_class.__name__ if dialect_class else None - if dialect_name in {"PGVector", "ParadeDB"}: + if isinstance(dialect, Postgres) or (dialect_name and dialect_name in _POSTGRES_EXTENSION_DIALECT_NAMES): if is_vector_distance_expression(expression): return render_vector_distance_postgres( generator.sql(expression, "this"), @@ -52,5 +58,6 @@ def _postgres_extension_operator_sql(generator: PostgresGenerator, expression: e invalidate_generator_dispatch(PostgresGenerator) +PGTextSearchGenerator = PostgresGenerator # pyright: ignore[reportAssignmentType] PGVectorGenerator = PostgresGenerator # pyright: ignore[reportAssignmentType] ParadeDBGenerator = PostgresGenerator # pyright: ignore[reportAssignmentType] diff --git a/sqlspec/dialects/postgres/_operators.py b/sqlspec/dialects/postgres/_operators.py index b8d97a8bd..4ea891a19 100644 --- a/sqlspec/dialects/postgres/_operators.py +++ b/sqlspec/dialects/postgres/_operators.py @@ -14,6 +14,7 @@ __all__ = ( "PARADEDB_OPERATOR_TOKENS", + "PG_TEXTSEARCH_OPERATOR_TOKENS", "PGVECTOR_OPERATOR_TOKENS", "is_postgres_extension_operator", "postgres_extension_operator", @@ -40,6 +41,9 @@ "##": TokenType.NESTED, "##>": TokenType.AGGREGATEFUNCTION, } +PG_TEXTSEARCH_OPERATOR_TOKENS: Final[dict[str, TokenType]] = { + "<@>": TokenType.RING, +} _REGISTERED = False @@ -61,7 +65,12 @@ def register_postgres_extension_operators() -> None: return factor: dict[TokenType, Any] = dict(PostgresParser.FACTOR) - for operator, token in {**PGVECTOR_OPERATOR_TOKENS, **PARADEDB_OPERATOR_TOKENS}.items(): + extension_tokens = { + **PGVECTOR_OPERATOR_TOKENS, + **PARADEDB_OPERATOR_TOKENS, + **PG_TEXTSEARCH_OPERATOR_TOKENS, + } + for operator, token in extension_tokens.items(): factor[token] = _build_operator_factory(operator) setattr(PostgresParser, "FACTOR", factor) diff --git a/sqlspec/dialects/postgres/_paradedb.py b/sqlspec/dialects/postgres/_paradedb.py index d6c813b00..db4411468 100644 --- a/sqlspec/dialects/postgres/_paradedb.py +++ b/sqlspec/dialects/postgres/_paradedb.py @@ -17,22 +17,31 @@ ``pyproject.toml`` and by the ``Dialect`` metaclass on import. """ +from sqlglot.dialects.postgres import Postgres + from sqlspec.dialects.postgres._generators import ParadeDBGenerator -from sqlspec.dialects.postgres._operators import PARADEDB_OPERATOR_TOKENS, register_postgres_extension_operators -from sqlspec.dialects.postgres._pgvector import PGVector, PGVectorTokenizer +from sqlspec.dialects.postgres._operators import ( + PARADEDB_OPERATOR_TOKENS, + PGVECTOR_OPERATOR_TOKENS, + register_postgres_extension_operators, +) __all__ = ("ParadeDB",) register_postgres_extension_operators() -class ParadeDBTokenizer(PGVectorTokenizer): +class ParadeDBTokenizer(Postgres.Tokenizer): """Tokenizer with ParadeDB search operators and pgvector distance operators.""" - KEYWORDS = {**PGVectorTokenizer.KEYWORDS, **PARADEDB_OPERATOR_TOKENS} + KEYWORDS = { + **Postgres.Tokenizer.KEYWORDS, + **PARADEDB_OPERATOR_TOKENS, + **PGVECTOR_OPERATOR_TOKENS, + } -class ParadeDB(PGVector): +class ParadeDB(Postgres): """ParadeDB dialect with pg_search and pgvector extension support.""" Tokenizer = ParadeDBTokenizer diff --git a/sqlspec/dialects/postgres/_pg_textsearch.py b/sqlspec/dialects/postgres/_pg_textsearch.py new file mode 100644 index 000000000..4baf8f8ec --- /dev/null +++ b/sqlspec/dialects/postgres/_pg_textsearch.py @@ -0,0 +1,42 @@ +"""PGTextSearch dialect extending Postgres with pg_textsearch BM25 operators. + +Adds support for PostgreSQL and AlloyDB pg_textsearch BM25 operators: + - <@> : Relevance ranking operator (returns negative BM25 score for ASC index scans) + +Scoring is handled via the <@> operator. Text search configuration and saturation +parameters (k1, b) are specified via index parameters in USING bm25. + +Also inherits the pgvector distance operators for seamless hybrid search. +Registered with sqlglot through the ``sqlglot.dialects`` entry-point group in +``pyproject.toml`` and by the ``Dialect`` metaclass on import. +""" + +from sqlglot.dialects.postgres import Postgres + +from sqlspec.dialects.postgres._generators import PGTextSearchGenerator +from sqlspec.dialects.postgres._operators import ( + PG_TEXTSEARCH_OPERATOR_TOKENS, + PGVECTOR_OPERATOR_TOKENS, + register_postgres_extension_operators, +) + +__all__ = ("PGTextSearch",) + +register_postgres_extension_operators() + + +class PGTextSearchTokenizer(Postgres.Tokenizer): + """Tokenizer with pg_textsearch BM25 ranking operators and pgvector distance operators.""" + + KEYWORDS = { + **Postgres.Tokenizer.KEYWORDS, + **PG_TEXTSEARCH_OPERATOR_TOKENS, + **PGVECTOR_OPERATOR_TOKENS, + } + + +class PGTextSearch(Postgres): + """PostgreSQL dialect with pg_textsearch and pgvector extension support.""" + + Tokenizer = PGTextSearchTokenizer + Generator = PGTextSearchGenerator diff --git a/tests/unit/dialects/test_operators.py b/tests/unit/dialects/test_operators.py new file mode 100644 index 000000000..89619568c --- /dev/null +++ b/tests/unit/dialects/test_operators.py @@ -0,0 +1,53 @@ +"""Unit tests for PostgreSQL extension operator tokens and FACTOR registration.""" + +from sqlglot import exp +from sqlglot.parsers.postgres import PostgresParser +from sqlglot.tokenizer_core import TokenType + +from sqlspec.dialects.postgres._operators import ( + PARADEDB_OPERATOR_TOKENS, + PG_TEXTSEARCH_OPERATOR_TOKENS, + PGVECTOR_OPERATOR_TOKENS, + is_postgres_extension_operator, + postgres_extension_operator, + register_postgres_extension_operators, +) + + +def test_pg_textsearch_operator_tokens_definition() -> None: + """Verify pg_textsearch operator token mapping.""" + assert "<@>" in PG_TEXTSEARCH_OPERATOR_TOKENS + assert PG_TEXTSEARCH_OPERATOR_TOKENS["<@>"] == TokenType.RING + + +def test_pg_textsearch_factor_registration() -> None: + """Verify pg_textsearch operator is registered in PostgresParser.FACTOR.""" + register_postgres_extension_operators() + ring_token = PG_TEXTSEARCH_OPERATOR_TOKENS["<@>"] + assert ring_token in PostgresParser.FACTOR + + factory = PostgresParser.FACTOR[ring_token] + left = exp.var("content") + right = exp.Literal.string("search query") + node = factory(left, right) + + assert isinstance(node, exp.Operator) + assert is_postgres_extension_operator(node) + assert postgres_extension_operator(node) == "<@>" + + +def test_no_operator_token_collisions() -> None: + """Verify extension operator sets are mutually disjoint.""" + all_operators = [ + *PGVECTOR_OPERATOR_TOKENS.keys(), + *PARADEDB_OPERATOR_TOKENS.keys(), + *PG_TEXTSEARCH_OPERATOR_TOKENS.keys(), + ] + assert len(all_operators) == len(set(all_operators)) + + all_tokens = [ + *PGVECTOR_OPERATOR_TOKENS.values(), + *PARADEDB_OPERATOR_TOKENS.values(), + *PG_TEXTSEARCH_OPERATOR_TOKENS.values(), + ] + assert len(all_tokens) == len(set(all_tokens)) diff --git a/tests/unit/dialects/test_paradedb.py b/tests/unit/dialects/test_paradedb.py index a9dec45b2..5f56c6e37 100644 --- a/tests/unit/dialects/test_paradedb.py +++ b/tests/unit/dialects/test_paradedb.py @@ -1,11 +1,11 @@ """Dialect unit tests for the ParadeDB (PostgreSQL + pgvector + pg_search) dialect.""" from sqlglot import parse_one +from sqlglot.dialects.postgres import Postgres import sqlspec.dialects.postgres._paradedb # noqa: F401 from sqlspec.dialects.postgres._operators import PARADEDB_OPERATOR_TOKENS, PGVECTOR_OPERATOR_TOKENS -from sqlspec.dialects.postgres._paradedb import ParadeDBTokenizer -from sqlspec.dialects.postgres._pgvector import PGVectorTokenizer +from sqlspec.dialects.postgres._paradedb import ParadeDB, ParadeDBTokenizer def _render(sql: str) -> str: @@ -101,8 +101,12 @@ def test_prox_regex() -> None: assert "@@@" in rendered -def test_paradedb_keywords_inherits_from_pgvector() -> None: - assert ParadeDBTokenizer.KEYWORDS == {**PGVectorTokenizer.KEYWORDS, **PARADEDB_OPERATOR_TOKENS} +def test_paradedb_keywords_inherits_from_postgres_and_extensions() -> None: + assert ParadeDBTokenizer.KEYWORDS == { + **Postgres.Tokenizer.KEYWORDS, + **PARADEDB_OPERATOR_TOKENS, + **PGVECTOR_OPERATOR_TOKENS, + } def test_paradedb_keywords_contains_paradedb_operators() -> None: @@ -115,5 +119,9 @@ def test_paradedb_keywords_contains_pgvector_operators() -> None: assert operator in ParadeDBTokenizer.KEYWORDS -def test_paradedb_tokenizer_inherits_from_pgvector_tokenizer() -> None: - assert issubclass(ParadeDBTokenizer, PGVectorTokenizer) +def test_paradedb_tokenizer_inherits_from_postgres_tokenizer() -> None: + assert issubclass(ParadeDBTokenizer, Postgres.Tokenizer) + + +def test_paradedb_subclasses_postgres() -> None: + assert issubclass(ParadeDB, Postgres) diff --git a/tests/unit/dialects/test_pg_textsearch.py b/tests/unit/dialects/test_pg_textsearch.py new file mode 100644 index 000000000..80b5b20b5 --- /dev/null +++ b/tests/unit/dialects/test_pg_textsearch.py @@ -0,0 +1,46 @@ +"""Unit tests for the PGTextSearch sqlglot dialect.""" + +from sqlglot import parse_one + +from sqlspec.dialects.postgres import PGTextSearch + + +def test_pg_textsearch_bm25_ranking_operator() -> None: + """Verify BM25 relevance ranking operator <@> parses and generates.""" + sql = "SELECT title, content <@> 'database system' AS score FROM documents ORDER BY content <@> 'database system' ASC" + expression = parse_one(sql, read=PGTextSearch) + rendered = expression.sql(dialect=PGTextSearch) + assert "<@>" in rendered + assert "ORDER BY content <@> 'database system' ASC" in rendered + + +def test_pg_textsearch_inherits_pgvector_distance() -> None: + """Verify pgvector distance operators are supported for hybrid search.""" + sql = "SELECT embedding <=> '[1,2,3]' FROM items" + expression = parse_one(sql, read=PGTextSearch) + rendered = expression.sql(dialect=PGTextSearch) + assert "<=>" in rendered + + +def test_pg_textsearch_hybrid_query() -> None: + """Verify combined vector distance and BM25 ranking in one statement.""" + sql = ( + "SELECT id, embedding <=> '[0.1, 0.2, 0.3]' AS vec_dist, content <@> 'fast database' AS bm25_score " + "FROM documents ORDER BY content <@> 'fast database' ASC LIMIT 10" + ) + expression = parse_one(sql, read=PGTextSearch) + rendered = expression.sql(dialect=PGTextSearch) + assert "<=>" in rendered + assert "<@>" in rendered + + +def test_pg_textsearch_bm25_index_ddl() -> None: + """Verify BM25 index creation statement parses.""" + sql = ( + "CREATE INDEX idx_docs_bm25 ON documents " + "USING bm25 (content) WITH (text_config='english', k1=1.2, b=0.75)" + ) + expression = parse_one(sql, read=PGTextSearch) + rendered = expression.sql(dialect=PGTextSearch) + assert "USING bm25" in rendered + assert "text_config" in rendered diff --git a/tests/unit/dialects/test_registration.py b/tests/unit/dialects/test_registration.py index 09541c8d0..cf9571d57 100644 --- a/tests/unit/dialects/test_registration.py +++ b/tests/unit/dialects/test_registration.py @@ -6,6 +6,7 @@ DIALECT_ENTRY_POINTS = { "paradedb": "sqlspec.dialects.postgres", + "pg_textsearch": "sqlspec.dialects.postgres", "pgvector": "sqlspec.dialects.postgres", "spangres": "sqlspec.dialects.spanner", "spanner": "sqlspec.dialects.spanner", @@ -52,13 +53,15 @@ def test_importing_sqlspec_does_not_eagerly_load_dialect_machinery() -> None: def test_lazy_dialects_attribute_still_works() -> None: code = ( "import sqlspec\n" - "from sqlspec.dialects import Spanner, Spangres, PGVector, ParadeDB\n" + "from sqlspec.dialects import Spanner, Spangres, PGVector, ParadeDB, PGTextSearch\n" "assert sqlspec.dialects.Spanner is Spanner\n" + "assert sqlspec.dialects.PGTextSearch is PGTextSearch\n" "from sqlglot.dialects.dialect import Dialect\n" "assert Dialect.get('spanner') is Spanner\n" "assert Dialect.get('spangres') is Spangres\n" "assert Dialect.get('pgvector') is PGVector\n" "assert Dialect.get('paradedb') is ParadeDB\n" + "assert Dialect.get('pg_textsearch') is PGTextSearch\n" "print('ok')\n" ) result = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True, check=False) From dfe3dbc47125fc5d112d1b2d1d935aaa81faed79 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 13 Sep 2026 17:14:47 +0000 Subject: [PATCH 02/10] feat(core): track active_extensions and probe pg_textsearch in config_runtime --- sqlspec/core/config_runtime.py | 29 ++++++++++++++++++++++++ tests/unit/core/test_config_runtime.py | 31 ++++++++++++++++++++++++++ tests/unit/core/test_splitter.py | 1 + 3 files changed, 61 insertions(+) create mode 100644 tests/unit/core/test_config_runtime.py diff --git a/sqlspec/core/config_runtime.py b/sqlspec/core/config_runtime.py index c28414b10..f8098282c 100644 --- a/sqlspec/core/config_runtime.py +++ b/sqlspec/core/config_runtime.py @@ -21,6 +21,7 @@ "close_sync_pool", "create_async_pool", "create_sync_pool", + "is_postgres_extension_active", "resolve_postgres_extension_state", "resolve_runtime_statement_config", "seed_runtime_driver_features", @@ -59,6 +60,8 @@ def build_postgres_extension_probe_names(driver_features: "dict[str, Any] | None extensions.append("vector") if driver_features.get("enable_paradedb", False): extensions.append("pg_search") + if driver_features.get("enable_pg_textsearch", False): + extensions.append("pg_textsearch") return extensions @@ -75,16 +78,42 @@ def resolve_postgres_extension_state( paradedb_available = bool( driver_features and driver_features.get("enable_paradedb", False) and "pg_search" in detected ) + pg_textsearch_available = bool( + driver_features and driver_features.get("enable_pg_textsearch", False) and "pg_textsearch" in detected + ) + + active_extensions: set[str] = set() + if pgvector_available: + active_extensions.add("vector") + if paradedb_available: + active_extensions.add("pg_search") + if pg_textsearch_available: + active_extensions.add("pg_textsearch") + + if driver_features is not None: + driver_features["active_extensions"] = active_extensions if statement_config.dialect == "postgres": if paradedb_available: statement_config = statement_config.replace(dialect="paradedb") + elif pg_textsearch_available: + statement_config = statement_config.replace(dialect="pg_textsearch") elif pgvector_available: statement_config = statement_config.replace(dialect="pgvector") return statement_config, pgvector_available, paradedb_available +def is_postgres_extension_active(driver_features: "dict[str, Any] | None", extension: str) -> bool: + """Return True if the named PostgreSQL extension is active in driver_features.""" + if driver_features is None: + return False + active = driver_features.get("active_extensions") + if isinstance(active, (set, frozenset, list, tuple)): + return extension in active + return False + + def resolve_runtime_statement_config( statement_config: StatementConfig | None, configured_statement_config: StatementConfig | None, diff --git a/tests/unit/core/test_config_runtime.py b/tests/unit/core/test_config_runtime.py new file mode 100644 index 000000000..34c66f4e5 --- /dev/null +++ b/tests/unit/core/test_config_runtime.py @@ -0,0 +1,31 @@ +"""Unit tests for sqlspec.core.config_runtime.""" + +from sqlspec.core.config_runtime import ( + build_postgres_extension_probe_names, + is_postgres_extension_active, + resolve_postgres_extension_state, +) +from sqlspec.core.statement import StatementConfig + + +def test_build_postgres_extension_probe_names_pg_textsearch() -> None: + features = {"enable_pgvector": True, "enable_paradedb": True, "enable_pg_textsearch": True} + probes = build_postgres_extension_probe_names(features) + assert probes == ["vector", "pg_search", "pg_textsearch"] + + +def test_resolve_postgres_extension_state_active_extensions() -> None: + features = {"enable_pgvector": True, "enable_pg_textsearch": True} + config = StatementConfig(dialect="postgres") + detected = {"vector", "pg_textsearch"} + + updated_config, pgvector_avail, paradedb_avail = resolve_postgres_extension_state(config, features, detected) + + assert pgvector_avail is True + assert paradedb_avail is False + assert updated_config.dialect == "pg_textsearch" + assert "active_extensions" in features + assert features["active_extensions"] == {"vector", "pg_textsearch"} + assert is_postgres_extension_active(features, "pg_textsearch") is True + assert is_postgres_extension_active(features, "vector") is True + assert is_postgres_extension_active(features, "pg_search") is False diff --git a/tests/unit/core/test_splitter.py b/tests/unit/core/test_splitter.py index c60c4a2b1..99b39d64e 100644 --- a/tests/unit/core/test_splitter.py +++ b/tests/unit/core/test_splitter.py @@ -45,6 +45,7 @@ def test_dialect_class_map_contains_expected_aliases() -> None: "postgres": splitter_module.PostgreSQLDialectConfig, "paradedb": splitter_module.PostgreSQLDialectConfig, "pgvector": splitter_module.PostgreSQLDialectConfig, + "pg_textsearch": splitter_module.PostgreSQLDialectConfig, "mysql": splitter_module.MySQLDialectConfig, "sqlite": splitter_module.SQLiteDialectConfig, "duckdb": splitter_module.DuckDBDialectConfig, From 31cc9e792874e33dad217802981a6e73b210e22f Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 13 Sep 2026 17:21:00 +0000 Subject: [PATCH 03/10] feat(adapters): add enable_pg_textsearch and standardize extension probing --- sqlspec/adapters/adbc/config.py | 28 +++++++++-- sqlspec/adapters/adbc/core.py | 22 +++++---- sqlspec/adapters/asyncpg/config.py | 24 +++++----- sqlspec/adapters/asyncpg/core.py | 3 ++ sqlspec/adapters/psqlpy/config.py | 13 +++++ sqlspec/adapters/psqlpy/core.py | 3 ++ sqlspec/adapters/psycopg/config.py | 42 ++++++++--------- sqlspec/adapters/psycopg/core.py | 3 ++ .../test_adbc/test_extension_detection.py | 47 +++++++++++++++---- .../unit/adapters/test_asyncpg/test_config.py | 41 +++++++++++----- .../unit/adapters/test_psqlpy/test_config.py | 20 +++++++- .../unit/adapters/test_psycopg/test_config.py | 13 +++++ 12 files changed, 192 insertions(+), 67 deletions(-) diff --git a/sqlspec/adapters/adbc/config.py b/sqlspec/adapters/adbc/config.py index 8e50b74a2..b5c957244 100644 --- a/sqlspec/adapters/adbc/config.py +++ b/sqlspec/adapters/adbc/config.py @@ -12,6 +12,7 @@ detect_postgres_extensions, get_statement_config, is_postgres_dialect, + is_postgres_extension_active, resolve_dialect_from_config, resolve_dialect_name, resolve_driver_connect_func, @@ -104,6 +105,10 @@ class AdbcDriverFeatures(TypedDict): When True and the resolved dialect is PostgreSQL, queries ``pg_extension`` on the first connection to check for the ``pg_search`` extension. Defaults to True. Independent of enable_pgvector. + enable_pg_textsearch: Enable pg_textsearch extension detection for BM25 search. + When True and the resolved dialect is PostgreSQL, queries ``pg_extension`` + on the first connection to check for the ``pg_textsearch`` extension. + Defaults to True. enable_events: Enable database event channel support. Defaults to True when extension_config["events"] is configured. Provides pub/sub capabilities via table-backed queue (ADBC has no native pub/sub). @@ -124,6 +129,7 @@ class AdbcDriverFeatures(TypedDict): arrow_extension_types: NotRequired[bool] enable_pgvector: NotRequired[bool] enable_paradedb: NotRequired[bool] + enable_pg_textsearch: NotRequired[bool] enable_events: NotRequired[bool] on_connection_create: "NotRequired[Callable[[AdbcConnection], None]]" events_backend: NotRequired[Literal["poll_queue"]] @@ -197,6 +203,7 @@ class AdbcConfig(NoPoolSyncConfig[AdbcConnection, AdbcDriver]): __slots__ = ( "_default_session_config", "_paradedb_available", + "_pg_textsearch_available", "_pgvector_available", "_resolved_dialect", "_user_connection_hook", @@ -231,6 +238,7 @@ def __init__( self.connection_config = normalize_connection_config(connection_config) self._pgvector_available: bool | None = None self._paradedb_available: bool | None = None + self._pg_textsearch_available: bool | None = None self._resolved_dialect = resolve_dialect_from_config(self.connection_config) @@ -288,7 +296,7 @@ def create_connection(self) -> AdbcConnection: def _update_dialect_for_extensions(self) -> None: """Update statement_config dialect based on detected extensions. - Priority: paradedb > pgvector > postgres (default). + Priority: paradedb > pg_textsearch > pgvector > postgres (default). Only switches when current dialect is ``postgres``. """ current_dialect = self.statement_config.dialect or "postgres" @@ -297,9 +305,16 @@ def _update_dialect_for_extensions(self) -> None: if self._paradedb_available: self.statement_config = self.statement_config.replace(dialect="paradedb") + elif self._pg_textsearch_available: + self.statement_config = self.statement_config.replace(dialect="pg_textsearch") elif self._pgvector_available: self.statement_config = self.statement_config.replace(dialect="pgvector") + @property + def pg_textsearch_available(self) -> bool: + """Return True if the pg_textsearch extension is available.""" + return bool(self._pg_textsearch_available) + def _detect_extensions_if_needed(self) -> None: """Detect postgres extensions on first call, caching results. @@ -313,13 +328,17 @@ def _detect_extensions_if_needed(self) -> None: if not is_postgres_dialect(dialect): self._pgvector_available = False self._paradedb_available = False + self._pg_textsearch_available = False return connection = self.create_connection() try: probe_names = build_postgres_extension_probe_names(self.driver_features) - pgvector_available, paradedb_available = detect_postgres_extensions( - connection, enable_pgvector="vector" in probe_names, enable_paradedb="pg_search" in probe_names + pgvector_available, paradedb_available, pg_textsearch_available = detect_postgres_extensions( + connection, + enable_pgvector="vector" in probe_names, + enable_paradedb="pg_search" in probe_names, + enable_pg_textsearch="pg_textsearch" in probe_names, ) finally: connection.close() @@ -329,9 +348,12 @@ def _detect_extensions_if_needed(self) -> None: detected_extensions.add("vector") if paradedb_available: detected_extensions.add("pg_search") + if pg_textsearch_available: + detected_extensions.add("pg_textsearch") self.statement_config, self._pgvector_available, self._paradedb_available = resolve_postgres_extension_state( self.statement_config, self.driver_features, detected_extensions ) + self._pg_textsearch_available = is_postgres_extension_active(self.driver_features, "pg_textsearch") def provide_connection(self, *args: Any, **kwargs: Any) -> "AdbcConnectionContext": """Provide a connection context manager. diff --git a/sqlspec/adapters/adbc/core.py b/sqlspec/adapters/adbc/core.py index 8e3c4ccbf..39e2f289f 100644 --- a/sqlspec/adapters/adbc/core.py +++ b/sqlspec/adapters/adbc/core.py @@ -21,6 +21,7 @@ ) from sqlspec.core.config_runtime import ( build_postgres_extension_probe_names, + is_postgres_extension_active, resolve_postgres_extension_state, resolve_runtime_statement_config, ) @@ -74,6 +75,7 @@ "get_statement_config", "handle_postgres_rollback", "is_postgres_dialect", + "is_postgres_extension_active", "normalize_driver_path", "normalize_postgres_empty_parameters", "normalize_script_rowcount", @@ -231,29 +233,32 @@ def detect_dialect(connection: Any, logger: Any | None = None, *, fallback_diale def detect_postgres_extensions( - connection: Any, *, enable_pgvector: bool = False, enable_paradedb: bool = False -) -> "tuple[bool, bool]": - """Detect pgvector and paradedb extensions on a postgres connection. + connection: Any, *, enable_pgvector: bool = False, enable_paradedb: bool = False, enable_pg_textsearch: bool = False +) -> "tuple[bool, bool, bool]": + """Detect pgvector, paradedb, and pg_textsearch extensions on a postgres connection. - Queries ``pg_extension`` for the ``vector`` and ``pg_search`` extensions. + Queries ``pg_extension`` for the ``vector``, ``pg_search``, and ``pg_textsearch`` extensions. Returns cached-friendly booleans suitable for storing on the config instance. Args: connection: ADBC connection to a PostgreSQL database. enable_pgvector: Whether to check for the pgvector extension. enable_paradedb: Whether to check for the pg_search extension. + enable_pg_textsearch: Whether to check for the pg_textsearch extension. Returns: - Tuple of ``(pgvector_available, paradedb_available)``. + Tuple of ``(pgvector_available, paradedb_available, pg_textsearch_available)``. """ extensions: list[str] = [] if enable_pgvector: extensions.append("vector") if enable_paradedb: extensions.append("pg_search") + if enable_pg_textsearch: + extensions.append("pg_textsearch") if not extensions: - return False, False + return False, False, False try: cursor = connection.cursor() @@ -261,11 +266,11 @@ def detect_postgres_extensions( cursor.execute("SELECT extname FROM pg_extension WHERE extname = ANY($1::text[])", [extensions]) rows = cursor.fetchall() detected: set[str] = {row[0] for row in rows} if rows else set() - return "vector" in detected, "pg_search" in detected + return "vector" in detected, "pg_search" in detected, "pg_textsearch" in detected finally: cursor.close() except Exception: - return False, False + return False, False, False def normalize_driver_path(driver_name: str) -> str: @@ -775,6 +780,7 @@ def apply_driver_features( processed_features.setdefault("enable_arrow_extension_types", processed_features["arrow_extension_types"]) processed_features.setdefault("enable_pgvector", PGVECTOR_INSTALLED) processed_features.setdefault("enable_paradedb", True) + processed_features.setdefault("enable_pg_textsearch", True) if json_serializer is not None: statement_config = _apply_adbc_json_serializer(statement_config, json_serializer) diff --git a/sqlspec/adapters/asyncpg/config.py b/sqlspec/adapters/asyncpg/config.py index 134c808df..ee0740d8f 100644 --- a/sqlspec/adapters/asyncpg/config.py +++ b/sqlspec/adapters/asyncpg/config.py @@ -21,6 +21,7 @@ build_connection_config, build_postgres_extension_probe_names, default_statement_config, + is_postgres_extension_active, register_json_codecs, register_pgvector_support, resolve_postgres_extension_state, @@ -131,6 +132,10 @@ class AsyncpgDriverFeatures(TypedDict): switches to "paradedb" which supports search operators (@@@, &&&, etc.) and inherits all pgvector distance operators. Defaults to True. Independent of enable_pgvector. + enable_pg_textsearch: Enable pg_textsearch extension detection for BM25 search. + When enabled and the pg_textsearch extension is detected, the SQL dialect + switches to "pg_textsearch" which supports the BM25 <@> relevance ranking operator. + Defaults to True. enable_cloud_sql: Enable Google Cloud SQL connector integration. Requires cloud-sql-python-connector package. Defaults to False (explicit opt-in required). @@ -175,6 +180,7 @@ class AsyncpgDriverFeatures(TypedDict): enable_json_codecs: NotRequired[bool] enable_pgvector: NotRequired[bool] enable_paradedb: NotRequired[bool] + enable_pg_textsearch: NotRequired[bool] enable_cloud_sql: NotRequired[bool] cloud_sql_instance: NotRequired[str] cloud_sql_enable_iam_auth: NotRequired[bool] @@ -454,14 +460,6 @@ async def _init_connection(self, connection: "AsyncpgConnection") -> None: if self._pgvector_available is None: detected_extensions: set[str] = set() extensions = build_postgres_extension_probe_names(self.driver_features) - adk_config = self.extension_config.get("adk", {}) - bm25_enabled = bool( - isinstance(adk_config, dict) - and adk_config.get("enable_memory", True) - and adk_config.get("enable_bm25", False) - ) - if bm25_enabled: - extensions.append("pg_textsearch") if extensions: try: results = await connection.fetch( @@ -470,12 +468,11 @@ async def _init_connection(self, connection: "AsyncpgConnection") -> None: detected_extensions = {r["extname"] for r in results} except Exception as exc: detected_extensions = set() - if bm25_enabled: - self._pg_textsearch_probe_error = exc + self._pg_textsearch_probe_error = exc self.statement_config, self._pgvector_available, self._paradedb_available = ( resolve_postgres_extension_state(self.statement_config, self.driver_features, detected_extensions) ) - self._pg_textsearch_available = "pg_textsearch" in detected_extensions if bm25_enabled else False + self._pg_textsearch_available = is_postgres_extension_active(self.driver_features, "pg_textsearch") if self._pgvector_available: await register_pgvector_support(connection) @@ -484,6 +481,11 @@ async def _init_connection(self, connection: "AsyncpgConnection") -> None: if self._user_connection_hook is not None: await self._user_connection_hook(connection) + @property + def pg_textsearch_available(self) -> bool: + """Return True if the pg_textsearch extension is available.""" + return bool(self._pg_textsearch_available) + def _ensure_pg_textsearch_available(self) -> None: if self._pg_textsearch_available: return diff --git a/sqlspec/adapters/asyncpg/core.py b/sqlspec/adapters/asyncpg/core.py index a541b5175..4e1f04c0f 100644 --- a/sqlspec/adapters/asyncpg/core.py +++ b/sqlspec/adapters/asyncpg/core.py @@ -11,6 +11,7 @@ from sqlspec.core import DriverParameterProfile, ParameterStyle, StatementConfig, build_statement_config_from_profile from sqlspec.core.config_runtime import ( build_postgres_extension_probe_names, + is_postgres_extension_active, resolve_postgres_extension_state, resolve_runtime_statement_config, ) @@ -58,6 +59,7 @@ "default_statement_config", "driver_profile", "invoke_prepared_statement", + "is_postgres_extension_active", "parse_status", "register_json_codecs", "register_pgvector_support", @@ -235,6 +237,7 @@ def apply_driver_features( processed_features.setdefault("enable_json_codecs", True) processed_features.setdefault("enable_pgvector", PGVECTOR_INSTALLED) processed_features.setdefault("enable_paradedb", True) + processed_features.setdefault("enable_pg_textsearch", True) processed_features.setdefault("enable_cloud_sql", False) processed_features.setdefault("enable_alloydb", False) diff --git a/sqlspec/adapters/psqlpy/config.py b/sqlspec/adapters/psqlpy/config.py index 2bce101d8..772cbbba8 100644 --- a/sqlspec/adapters/psqlpy/config.py +++ b/sqlspec/adapters/psqlpy/config.py @@ -11,6 +11,7 @@ build_connection_config, build_postgres_extension_probe_names, default_statement_config, + is_postgres_extension_active, resolve_postgres_extension_state, resolve_runtime_statement_config, ) @@ -100,6 +101,10 @@ class PsqlpyDriverFeatures(TypedDict): switches to "paradedb" which supports search operators (@@@, &&&, etc.) and inherits all pgvector distance operators. Defaults to True. Independent of enable_pgvector. + enable_pg_textsearch: Enable pg_textsearch extension detection for BM25 search. + When enabled and the pg_textsearch extension is detected, the SQL dialect + switches to "pg_textsearch" which supports the BM25 <@> relevance ranking operator. + Defaults to True. json_serializer: Custom JSON serializer applied to the statement configuration. json_deserializer: Custom JSON deserializer retained alongside the serializer for parity with asyncpg. on_connection_create: Async callback executed when a connection is acquired from pool. @@ -120,6 +125,7 @@ class PsqlpyDriverFeatures(TypedDict): enable_cast_detection: NotRequired[bool] enable_pgvector: NotRequired[bool] enable_paradedb: NotRequired[bool] + enable_pg_textsearch: NotRequired[bool] json_serializer: NotRequired["Callable[[Any], str]"] json_deserializer: NotRequired["Callable[[str], Any]"] on_connection_create: "NotRequired[Callable[[PsqlpyConnection], Awaitable[None]]]" @@ -237,6 +243,7 @@ def __init__( self._initialized_connection_ids: set[int] = set() self._pgvector_available: bool | None = None self._paradedb_available: bool | None = None + self._pg_textsearch_available: bool | None = None super().__init__( connection_config=connection_config, @@ -249,6 +256,11 @@ def __init__( **kwargs, ) + @property + def pg_textsearch_available(self) -> bool: + """Return True if the pg_textsearch extension is available.""" + return bool(self._pg_textsearch_available) + async def _ensure_connection(self, connection: "PsqlpyConnection") -> None: """Ensure connection callback has been called exactly once for this connection. @@ -269,6 +281,7 @@ async def _ensure_connection(self, connection: "PsqlpyConnection") -> None: self.statement_config, self._pgvector_available, self._paradedb_available = ( resolve_postgres_extension_state(self.statement_config, self.driver_features, detected_extensions) ) + self._pg_textsearch_available = is_postgres_extension_active(self.driver_features, "pg_textsearch") conn_id = id(connection) if conn_id in self._initialized_connection_ids: diff --git a/sqlspec/adapters/psqlpy/core.py b/sqlspec/adapters/psqlpy/core.py index bfe9dcc77..a598aee76 100644 --- a/sqlspec/adapters/psqlpy/core.py +++ b/sqlspec/adapters/psqlpy/core.py @@ -23,6 +23,7 @@ ) from sqlspec.core.config_runtime import ( build_postgres_extension_probe_names, + is_postgres_extension_active, resolve_postgres_extension_state, resolve_runtime_statement_config, ) @@ -74,6 +75,7 @@ "format_execute_many_parameters", "format_table_identifier", "get_parameter_casts", + "is_postgres_extension_active", "prepare_parameters_with_casts", "resolve_postgres_extension_state", "resolve_runtime_statement_config", @@ -171,6 +173,7 @@ def apply_driver_features( features.setdefault("enable_cast_detection", True) features.setdefault("enable_pgvector", PGVECTOR_INSTALLED) features.setdefault("enable_paradedb", True) + features.setdefault("enable_pg_textsearch", True) base_config = build_statement_config_from_profile( driver_profile, json_serializer=serializer, json_deserializer=deserializer diff --git a/sqlspec/adapters/psycopg/config.py b/sqlspec/adapters/psycopg/config.py index b045beefc..006599aa1 100644 --- a/sqlspec/adapters/psycopg/config.py +++ b/sqlspec/adapters/psycopg/config.py @@ -19,6 +19,7 @@ apply_driver_features, build_postgres_extension_probe_names, default_statement_config, + is_postgres_extension_active, resolve_postgres_extension_state, resolve_runtime_statement_config, ) @@ -124,6 +125,10 @@ class PsycopgDriverFeatures(TypedDict): switches to "paradedb" which supports search operators (@@@, &&&, etc.) and inherits all pgvector distance operators. Defaults to True. Independent of enable_pgvector. + enable_pg_textsearch: Enable pg_textsearch extension detection for BM25 search. + When enabled and the pg_textsearch extension is detected, the SQL dialect + switches to "pg_textsearch" which supports the BM25 <@> relevance ranking operator. + Defaults to True. json_serializer: Custom JSON serializer for StatementConfig parameter handling. json_deserializer: Custom JSON deserializer reference stored alongside the serializer for parity with asyncpg. on_connection_create: Callback executed when a connection is created/acquired from the pool. @@ -152,6 +157,7 @@ class PsycopgDriverFeatures(TypedDict): enable_pgvector: NotRequired[bool] enable_paradedb: NotRequired[bool] + enable_pg_textsearch: NotRequired[bool] json_serializer: NotRequired["Callable[[Any], str]"] json_deserializer: NotRequired["Callable[[str], Any]"] on_connection_create: NotRequired["Callable[..., Any]"] @@ -430,14 +436,6 @@ def _configure_connection(self, conn: "PsycopgSyncConnection") -> None: if self._pgvector_available is None: detected_extensions: set[str] = set() extensions = build_postgres_extension_probe_names(self.driver_features) - adk_config = self.extension_config.get("adk", {}) - bm25_enabled = bool( - isinstance(adk_config, dict) - and adk_config.get("enable_memory", True) - and adk_config.get("enable_bm25", False) - ) - if bm25_enabled: - extensions.append("pg_textsearch") if extensions: try: cursor = conn.execute( @@ -447,12 +445,11 @@ def _configure_connection(self, conn: "PsycopgSyncConnection") -> None: detected_extensions = {r[0] for r in results} # type: ignore[index] except Exception as exc: detected_extensions = set() - if bm25_enabled: - self._pg_textsearch_probe_error = exc + self._pg_textsearch_probe_error = exc self.statement_config, self._pgvector_available, self._paradedb_available = ( resolve_postgres_extension_state(self.statement_config, self.driver_features, detected_extensions) ) - self._pg_textsearch_available = "pg_textsearch" in detected_extensions if bm25_enabled else False + self._pg_textsearch_available = is_postgres_extension_active(self.driver_features, "pg_textsearch") if self._pgvector_available: register_pgvector_sync(conn) @@ -465,6 +462,11 @@ def _configure_connection(self, conn: "PsycopgSyncConnection") -> None: if self._user_connection_hook is not None: self._user_connection_hook(conn) + @property + def pg_textsearch_available(self) -> bool: + """Return True if the pg_textsearch extension is available.""" + return bool(self._pg_textsearch_available) + def _ensure_pg_textsearch_available(self) -> None: if self._pg_textsearch_available: return @@ -721,14 +723,6 @@ async def _configure_async_connection(self, conn: "PsycopgAsyncConnection") -> N if self._pgvector_available is None: detected_extensions: set[str] = set() extensions = build_postgres_extension_probe_names(self.driver_features) - adk_config = self.extension_config.get("adk", {}) - bm25_enabled = bool( - isinstance(adk_config, dict) - and adk_config.get("enable_memory", True) - and adk_config.get("enable_bm25", False) - ) - if bm25_enabled: - extensions.append("pg_textsearch") if extensions: try: cursor = await conn.execute( @@ -738,12 +732,11 @@ async def _configure_async_connection(self, conn: "PsycopgAsyncConnection") -> N detected_extensions = {r[0] for r in results} # type: ignore[index] except Exception as exc: detected_extensions = set() - if bm25_enabled: - self._pg_textsearch_probe_error = exc + self._pg_textsearch_probe_error = exc self.statement_config, self._pgvector_available, self._paradedb_available = ( resolve_postgres_extension_state(self.statement_config, self.driver_features, detected_extensions) ) - self._pg_textsearch_available = "pg_textsearch" in detected_extensions if bm25_enabled else False + self._pg_textsearch_available = is_postgres_extension_active(self.driver_features, "pg_textsearch") if self._pgvector_available: await register_pgvector_async(conn) @@ -756,6 +749,11 @@ async def _configure_async_connection(self, conn: "PsycopgAsyncConnection") -> N if self._user_connection_hook is not None: await self._user_connection_hook(conn) + @property + def pg_textsearch_available(self) -> bool: + """Return True if the pg_textsearch extension is available.""" + return bool(self._pg_textsearch_available) + def _ensure_pg_textsearch_available(self) -> None: if self._pg_textsearch_available: return diff --git a/sqlspec/adapters/psycopg/core.py b/sqlspec/adapters/psycopg/core.py index 2a98b5c96..04cc12497 100644 --- a/sqlspec/adapters/psycopg/core.py +++ b/sqlspec/adapters/psycopg/core.py @@ -17,6 +17,7 @@ ) from sqlspec.core.config_runtime import ( build_postgres_extension_probe_names, + is_postgres_extension_active, resolve_postgres_extension_state, resolve_runtime_statement_config, ) @@ -75,6 +76,7 @@ "driver_profile", "execute_with_optional_parameters", "execute_with_optional_parameters_async", + "is_postgres_extension_active", "pipeline_supported", "resolve_many_rowcount", "resolve_postgres_extension_state", @@ -175,6 +177,7 @@ def apply_driver_features( features.setdefault("json_deserializer", deserializer) features.setdefault("enable_pgvector", PGVECTOR_INSTALLED) features.setdefault("enable_paradedb", True) + features.setdefault("enable_pg_textsearch", True) parameter_config = _parameter_config(driver_profile, serializer, deserializer) statement_config = statement_config.replace(parameter_config=parameter_config) diff --git a/tests/unit/adapters/test_adbc/test_extension_detection.py b/tests/unit/adapters/test_adbc/test_extension_detection.py index 921e4f036..bd39984cd 100644 --- a/tests/unit/adapters/test_adbc/test_extension_detection.py +++ b/tests/unit/adapters/test_adbc/test_extension_detection.py @@ -70,51 +70,65 @@ def test_build_postgres_extension_probe_names_filters_disabled_features() -> Non def test_detect_postgres_extensions_returns_tuple() -> None: - """detect_postgres_extensions returns (pgvector_available, paradedb_available).""" + """detect_postgres_extensions returns (pgvector_available, paradedb_available, pg_textsearch_available).""" cursor = _Cursor([("vector",)]) connection = _Connection(cursor) - pgvector, paradedb = detect_postgres_extensions(connection, enable_pgvector=True, enable_paradedb=True) + pgvector, paradedb, pg_textsearch = detect_postgres_extensions( + connection, enable_pgvector=True, enable_paradedb=True, enable_pg_textsearch=True + ) assert pgvector is True assert paradedb is False + assert pg_textsearch is False assert cursor.closed is True def test_detect_postgres_extensions_both_available() -> None: - """Both extensions detected when both present.""" - connection = _Connection(_Cursor([("vector",), ("pg_search",)])) + """All extensions detected when present.""" + connection = _Connection(_Cursor([("vector",), ("pg_search",), ("pg_textsearch",)])) - pgvector, paradedb = detect_postgres_extensions(connection, enable_pgvector=True, enable_paradedb=True) + pgvector, paradedb, pg_textsearch = detect_postgres_extensions( + connection, enable_pgvector=True, enable_paradedb=True, enable_pg_textsearch=True + ) assert pgvector is True assert paradedb is True + assert pg_textsearch is True def test_detect_postgres_extensions_none_enabled() -> None: - """Returns (False, False) when both flags disabled.""" + """Returns (False, False, False) when all flags disabled.""" connection = _Connection(_Cursor([])) - pgvector, paradedb = detect_postgres_extensions(connection, enable_pgvector=False, enable_paradedb=False) + pgvector, paradedb, pg_textsearch = detect_postgres_extensions( + connection, enable_pgvector=False, enable_paradedb=False, enable_pg_textsearch=False + ) assert pgvector is False assert paradedb is False + assert pg_textsearch is False assert connection.cursor_requested is False def test_detect_postgres_extensions_handles_error() -> None: - """Returns (False, False) on query failure.""" + """Returns (False, False, False) on query failure.""" cursor = _Cursor([], error=Exception("connection error")) connection = _Connection(cursor) - pgvector, paradedb = detect_postgres_extensions(connection, enable_pgvector=True, enable_paradedb=True) + pgvector, paradedb, pg_textsearch = detect_postgres_extensions( + connection, enable_pgvector=True, enable_paradedb=True, enable_pg_textsearch=True + ) assert pgvector is False assert paradedb is False + assert pg_textsearch is False assert cursor.closed is True def test_adbc_config_initializes_extension_flags_to_none() -> None: - """AdbcConfig starts with _pgvector_available and _paradedb_available as None.""" + """AdbcConfig starts with extension flags as None.""" config = AdbcConfig(connection_config={"uri": ":memory:", "driver_name": "sqlite"}) assert config._pgvector_available is None # pyright: ignore[reportPrivateUsage] assert config._paradedb_available is None # pyright: ignore[reportPrivateUsage] + assert config._pg_textsearch_available is None # pyright: ignore[reportPrivateUsage] + assert config.pg_textsearch_available is False def test_resolve_postgres_extension_state_promotes_paradedb() -> None: @@ -133,15 +147,28 @@ def test_adbc_config_update_dialect_for_extensions_pgvector() -> None: config = AdbcConfig(connection_config={"uri": "postgresql://localhost/test"}) config._pgvector_available = True # pyright: ignore[reportPrivateUsage] config._paradedb_available = False # pyright: ignore[reportPrivateUsage] + config._pg_textsearch_available = False # pyright: ignore[reportPrivateUsage] config._update_dialect_for_extensions() # pyright: ignore[reportPrivateUsage] assert config.statement_config.dialect == "pgvector" +def test_adbc_config_update_dialect_for_extensions_pg_textsearch() -> None: + """Dialect switches to pg_textsearch when pg_textsearch is available.""" + config = AdbcConfig(connection_config={"uri": "postgresql://localhost/test"}) + config._pgvector_available = True # pyright: ignore[reportPrivateUsage] + config._paradedb_available = False # pyright: ignore[reportPrivateUsage] + config._pg_textsearch_available = True # pyright: ignore[reportPrivateUsage] + config._update_dialect_for_extensions() # pyright: ignore[reportPrivateUsage] + assert config.statement_config.dialect == "pg_textsearch" + assert config.pg_textsearch_available is True + + def test_adbc_config_update_dialect_for_extensions_paradedb() -> None: """Dialect switches to paradedb when both extensions available (paradedb > pgvector).""" config = AdbcConfig(connection_config={"uri": "postgresql://localhost/test"}) config._pgvector_available = True # pyright: ignore[reportPrivateUsage] config._paradedb_available = True # pyright: ignore[reportPrivateUsage] + config._pg_textsearch_available = True # pyright: ignore[reportPrivateUsage] config._update_dialect_for_extensions() # pyright: ignore[reportPrivateUsage] assert config.statement_config.dialect == "paradedb" diff --git a/tests/unit/adapters/test_asyncpg/test_config.py b/tests/unit/adapters/test_asyncpg/test_config.py index 1bfc7467b..b0154a88e 100644 --- a/tests/unit/adapters/test_asyncpg/test_config.py +++ b/tests/unit/adapters/test_asyncpg/test_config.py @@ -206,18 +206,35 @@ async def test_asyncpg_bm25_probe_failure_is_preserved_as_cause() -> None: assert exc_info.value.__cause__ is probe_error -async def test_asyncpg_does_not_probe_pg_textsearch_when_memory_or_bm25_is_disabled() -> None: - """Unneeded BM25 capability checks are omitted from first-connection setup.""" - for adk_config in ({"enable_memory": False, "enable_bm25": True}, {"enable_bm25": False}): - connection = AsyncMock() - config = AsyncpgConfig( - driver_features={"enable_json_codecs": False, "enable_pgvector": False, "enable_paradedb": False}, - extension_config={"adk": adk_config}, - ) - - await config._init_connection(connection) # pyright: ignore[reportPrivateUsage] - - connection.fetch.assert_not_awaited() +async def test_asyncpg_does_not_probe_pg_textsearch_when_disabled() -> None: + """Extension probe is omitted when all extension flags are disabled.""" + connection = AsyncMock() + config = AsyncpgConfig( + driver_features={ + "enable_json_codecs": False, + "enable_pgvector": False, + "enable_paradedb": False, + "enable_pg_textsearch": False, + } + ) + + await config._init_connection(connection) # pyright: ignore[reportPrivateUsage] + + connection.fetch.assert_not_awaited() + + +async def test_asyncpg_pg_textsearch_available_property() -> None: + """pg_textsearch_available reflects detected extension state.""" + connection = AsyncMock() + connection.fetch.return_value = [{"extname": "pg_textsearch"}] + config = AsyncpgConfig( + driver_features={"enable_json_codecs": False, "enable_pgvector": False, "enable_paradedb": False} + ) + + assert config.pg_textsearch_available is False + await config._init_connection(connection) # pyright: ignore[reportPrivateUsage] + assert config.pg_textsearch_available is True + assert config.statement_config.dialect == "pg_textsearch" @pytest.mark.anyio diff --git a/tests/unit/adapters/test_psqlpy/test_config.py b/tests/unit/adapters/test_psqlpy/test_config.py index af71eb21a..2670ce359 100644 --- a/tests/unit/adapters/test_psqlpy/test_config.py +++ b/tests/unit/adapters/test_psqlpy/test_config.py @@ -135,7 +135,9 @@ def test_psqlpy_resolve_postgres_extension_state_promotes_paradedb() -> None: @pytest.mark.anyio async def test_psqlpy_enable_pgvector_detects_extension_and_promotes_dialect() -> None: """Psqlpy pgvector support should mean extension detection and dialect promotion.""" - config = PsqlpyConfig(driver_features={"enable_pgvector": True, "enable_paradedb": False}) + config = PsqlpyConfig( + driver_features={"enable_pgvector": True, "enable_paradedb": False, "enable_pg_textsearch": False} + ) connection = _ExtensionConnection({"vector"}) await config._ensure_connection(cast("PsqlpyConnection", connection)) # pyright: ignore[reportPrivateUsage] @@ -145,6 +147,22 @@ async def test_psqlpy_enable_pgvector_detects_extension_and_promotes_dialect() - assert config.statement_config.dialect == "pgvector" +@pytest.mark.anyio +async def test_psqlpy_enable_pg_textsearch_detects_extension_and_promotes_dialect() -> None: + """Psqlpy pg_textsearch support should mean extension detection and dialect promotion.""" + config = PsqlpyConfig( + driver_features={"enable_pgvector": False, "enable_paradedb": False, "enable_pg_textsearch": True} + ) + connection = _ExtensionConnection({"pg_textsearch"}) + + assert config.pg_textsearch_available is False + await config._ensure_connection(cast("PsqlpyConnection", connection)) # pyright: ignore[reportPrivateUsage] + + assert connection.queries[0][1] == [["pg_textsearch"]] + assert config.pg_textsearch_available is True + assert config.statement_config.dialect == "pg_textsearch" + + @pytest.mark.anyio async def test_psqlpy_session_context_resolves_callable_statement_config() -> None: """Session context should call statement_config when it's a callable.""" diff --git a/tests/unit/adapters/test_psycopg/test_config.py b/tests/unit/adapters/test_psycopg/test_config.py index b340ac64f..29a9076e3 100644 --- a/tests/unit/adapters/test_psycopg/test_config.py +++ b/tests/unit/adapters/test_psycopg/test_config.py @@ -180,6 +180,19 @@ async def test_psycopg_async_bm25_probe_failure_is_preserved_without_callback_fa assert exc_info.value.__cause__ is probe_error +def test_psycopg_pg_textsearch_available_property() -> None: + """pg_textsearch_available reflects detected extension state.""" + connection = MagicMock() + connection.autocommit = True + connection.execute.return_value.fetchall.return_value = [("pg_textsearch",)] + config = PsycopgSyncConfig(driver_features={"enable_pgvector": False, "enable_paradedb": False}) + + assert config.pg_textsearch_available is False + config._configure_connection(connection) # pyright: ignore[reportPrivateUsage] + assert config.pg_textsearch_available is True + assert config.statement_config.dialect == "pg_textsearch" + + def test_psycopg_numeric_placeholders_convert_to_pyformat() -> None: """Numeric placeholders should be rewritten for psycopg execution.""" From b3536467bf6d5c9cd46ae746ca975d7a28091be0 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 13 Sep 2026 17:24:44 +0000 Subject: [PATCH 04/10] docs(postgres): document pg_textsearch dialect and register in mypyc inventory --- docs/changelog.rst | 22 ++++++++++++++++++++++ docs/reference/adapters/adbc.rst | 12 +++++++----- docs/reference/adapters/asyncpg.rst | 2 +- docs/reference/adapters/psqlpy.rst | 2 +- docs/reference/adapters/psycopg.rst | 2 +- docs/reference/dialects.rst | 24 ++++++++++++++++++++++-- pyproject.toml | 1 + tools/scripts/mypyc_boundary_map.py | 4 ++++ tools/scripts/mypyc_inventory.py | 5 +++++ 9 files changed, 64 insertions(+), 10 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 078016b28..a3d9d6f96 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -9,6 +9,28 @@ important operational fixes. Recent Updates ============== +Unreleased +---------- + +**Added:** + +* Native AlloyDB and PostgreSQL BM25 full-text search support via the ``pg_textsearch`` extension. + Includes the ``PGTextSearch`` dialect registered under ``sqlglot.dialects``, custom AST operator + support for the BM25 relevance ranking operator (``<@>``), automatic extension detection, and + ``enable_pg_textsearch`` configuration across all PostgreSQL adapters (AsyncPG, Psycopg, ADBC, and PsqlPy). +* Public :func:`~sqlspec.core.config_runtime.is_postgres_extension_active` helper in ``sqlspec.core.config_runtime`` + and adapter modules, with ``active_extensions`` capability tracking on runtime driver features. +* Exposed ``pg_textsearch_available`` property across ``AsyncpgConfig``, ``PsycopgSyncConfig``, + ``PsycopgAsyncConfig``, ``AdbcConfig``, and ``PsqlpyConfig``. + +**Changed:** + +* Decoupled the ``ParadeDB`` dialect so it inherits directly from ``Postgres`` rather than ``PGVector``, + allowing clean independent combinations of vector search and BM25 extensions. +* Standardized PostgreSQL extension detection across all adapters on a single first-connection probe + via :func:`~sqlspec.core.config_runtime.build_postgres_extension_probe_names`, removing ad-hoc ADK + probe branches. + v0.63.0 - Transactions, table fixtures, SQL fragments, storage, and kwargs parameter binding --------------------------------------------------------------------------------------------------- diff --git a/docs/reference/adapters/adbc.rst b/docs/reference/adapters/adbc.rst index c4743851e..a73b4d309 100644 --- a/docs/reference/adapters/adbc.rst +++ b/docs/reference/adapters/adbc.rst @@ -29,7 +29,7 @@ Supported Backends * - PostgreSQL - ``postgres``, ``postgresql``, ``pg`` - ``adbc_driver_postgresql`` - - Numeric parameters; optional pgvector and ParadeDB dialect detection. + - Numeric parameters; optional pgvector, pg_textsearch, and ParadeDB dialect detection. * - SQLite - ``sqlite``, ``sqlite3`` - ``adbc_driver_sqlite`` @@ -212,13 +212,15 @@ first connection and upgrades the SQL dialect accordingly: - **pgvector** — If the ``vector`` extension is installed, switches to the ``pgvector`` dialect which supports distance operators (``<->``, ``<=>``, ``<#>``, ``<+>``, ``<~>``, ``<%>``). -- **ParadeDB** — If the ``pg_search`` extension is installed (alongside ``vector``), - switches to the ``paradedb`` dialect which adds BM25 search operators (``@@@``, ``&&&``, - ``|||``, ``===``) on top of pgvector operators. +- **pg_textsearch** — If the ``pg_textsearch`` extension is installed, switches to the + ``pg_textsearch`` dialect which supports BM25 score ranking (``<@>``). +- **ParadeDB** — If the ``pg_search`` extension is installed, switches to the ``paradedb`` + dialect which adds BM25 search operators (``@@@``, ``&&&``, ``|||``, ``===``). -Detection is controlled by two driver feature flags: +Detection is controlled by driver feature flags: - ``enable_pgvector`` — Defaults to ``True`` when the ``pgvector`` Python package is installed. +- ``enable_pg_textsearch`` — Defaults to ``True``. - ``enable_paradedb`` — Defaults to ``True``. Detection runs once per config instance and caches the result. Non-PostgreSQL backends diff --git a/docs/reference/adapters/asyncpg.rst b/docs/reference/adapters/asyncpg.rst index 3016fb936..241b5571a 100644 --- a/docs/reference/adapters/asyncpg.rst +++ b/docs/reference/adapters/asyncpg.rst @@ -53,7 +53,7 @@ Driver Extension Dialects ================== -AsyncPG supports the :doc:`pgvector and ParadeDB dialects <../dialects>` for vector +AsyncPG supports the :doc:`pgvector, pg_textsearch, and ParadeDB dialects <../dialects>` for vector similarity search and full-text search operators. See the :doc:`Dialects <../dialects>` reference for operator details. diff --git a/docs/reference/adapters/psqlpy.rst b/docs/reference/adapters/psqlpy.rst index 4faa486e3..0fbf489f2 100644 --- a/docs/reference/adapters/psqlpy.rst +++ b/docs/reference/adapters/psqlpy.rst @@ -36,7 +36,7 @@ Driver Features Extension Dialects ================== -PsqlPy supports the :doc:`pgvector and ParadeDB dialects <../dialects>` for vector +PsqlPy supports the :doc:`pgvector, pg_textsearch, and ParadeDB dialects <../dialects>` for vector similarity search and full-text search operators. See the :doc:`Dialects <../dialects>` reference for operator details. diff --git a/docs/reference/adapters/psycopg.rst b/docs/reference/adapters/psycopg.rst index 2ee8b746f..8be287ccd 100644 --- a/docs/reference/adapters/psycopg.rst +++ b/docs/reference/adapters/psycopg.rst @@ -78,7 +78,7 @@ Async Driver Extension Dialects ================== -Psycopg supports the :doc:`pgvector and ParadeDB dialects <../dialects>` for vector +Psycopg supports the :doc:`pgvector, pg_textsearch, and ParadeDB dialects <../dialects>` for vector similarity search and full-text search operators. See the :doc:`Dialects <../dialects>` reference for operator details. diff --git a/docs/reference/dialects.rst b/docs/reference/dialects.rst index 8c64b39e7..3bdb61aa8 100644 --- a/docs/reference/dialects.rst +++ b/docs/reference/dialects.rst @@ -12,7 +12,7 @@ group, so ``sqlglot.parse_one(..., dialect="pgvector")`` resolves them in any environment where SQLSpec is installed; importing ``sqlspec`` alone does not load them. Import the classes directly when you need them as objects:: - from sqlspec.dialects import PGVector, ParadeDB, Spanner, Spangres + from sqlspec.dialects import PGTextSearch, PGVector, ParadeDB, Spanner, Spangres Performance builds compile the custom dialect helper modules alongside ``sqlglot[c]``: generator transforms, operator registries, and compatibility @@ -50,6 +50,26 @@ Adds support for pgvector distance operators: * - ``<%>`` - Jaccard distance (binary vectors) +PGTextSearch +------------ + +.. autoclass:: sqlspec.dialects.postgres.PGTextSearch + :members: + :show-inheritance: + :no-index: + +Adds support for the native PostgreSQL and Google Cloud AlloyDB / AlloyDB Omni ``pg_textsearch`` BM25 extension: + +.. list-table:: + :header-rows: 1 + + * - Operator + - Description + * - ``<@>`` + - BM25 score ranking operator (returns negative score for ASC index scans) + +Indices are created with ``USING bm25 (column) WITH (text_config='english')`` and queries order by ``column <@> 'query' ASC``. + ParadeDB -------- @@ -58,7 +78,7 @@ ParadeDB :show-inheritance: :no-index: -Extends PGVector with ParadeDB pg_search operators: +Extends PostgreSQL with ParadeDB (pg_search) operators: .. list-table:: :header-rows: 1 diff --git a/pyproject.toml b/pyproject.toml index 9ad2125de..bb1aad4cc 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -237,6 +237,7 @@ exclude = [ "sqlspec/migrations/commands.py", # interpreted: inspect.signature/functools.wraps on coroutines "sqlspec/dialects/postgres/_pgvector.py", # interpreted: Dialect subclass needs sqlglot's metaclass "sqlspec/dialects/postgres/_paradedb.py", # interpreted: Dialect subclass needs sqlglot's metaclass + "sqlspec/dialects/postgres/_pg_textsearch.py", # interpreted: Dialect subclass needs sqlglot's metaclass "sqlspec/dialects/spanner/_spanner.py", # interpreted: Dialect subclass needs sqlglot's metaclass "sqlspec/dialects/spanner/_spangres.py", # interpreted: Dialect subclass needs sqlglot's metaclass "sqlspec/storage/_arrow_payload.py", # PyArrow conversion boundary stays interpreted diff --git a/tools/scripts/mypyc_boundary_map.py b/tools/scripts/mypyc_boundary_map.py index aef355c74..9903b85e5 100644 --- a/tools/scripts/mypyc_boundary_map.py +++ b/tools/scripts/mypyc_boundary_map.py @@ -171,6 +171,10 @@ "bucket": "hard_block", "reason": "SQLGlot dialect subclass module fails native class import under mypyc.", }, + "sqlspec/dialects/postgres/_pg_textsearch.py": { + "bucket": "hard_block", + "reason": "SQLGlot dialect subclass module fails native class import under mypyc.", + }, "sqlspec/dialects/spanner/_spanner.py": { "bucket": "hard_block", "reason": "SQLGlot tokenizer/dialect subclass module fails native class import under mypyc.", diff --git a/tools/scripts/mypyc_inventory.py b/tools/scripts/mypyc_inventory.py index 814490eb5..5049c93f8 100644 --- a/tools/scripts/mypyc_inventory.py +++ b/tools/scripts/mypyc_inventory.py @@ -109,6 +109,10 @@ "classification": "hard_block", "reason": "SQLGlot subclass/registration module fails native class import under mypyc; compiled helpers stay in _generators/_operators.", }, + "sqlspec/dialects/postgres/_pg_textsearch.py": { + "classification": "hard_block", + "reason": "SQLGlot subclass/registration module fails native class import under mypyc; compiled helpers stay in _generators/_operators.", + }, "sqlspec/dialects/postgres/_pgvector.py": { "classification": "hard_block", "reason": "SQLGlot tokenizer/dialect subclass module fails native class import under mypyc; compiled helpers stay in _generators/_operators.", @@ -407,6 +411,7 @@ def build_inventory(root: Path | None = None) -> dict[str, Any]: if pattern in { "sqlspec/dialects/postgres/_paradedb.py", + "sqlspec/dialects/postgres/_pg_textsearch.py", "sqlspec/dialects/postgres/_pgvector.py", "sqlspec/dialects/spanner/_spangres.py", "sqlspec/dialects/spanner/_spanner.py", From cf92fe36dec956326163389bb0614a0329d7f118 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 13 Sep 2026 17:26:01 +0000 Subject: [PATCH 05/10] style(dialects): format dialect files with ruff --- sqlspec/dialects/postgres/_generators.py | 9 ++++++--- sqlspec/dialects/postgres/_operators.py | 12 +++--------- sqlspec/dialects/postgres/_paradedb.py | 6 +----- sqlspec/dialects/postgres/_pg_textsearch.py | 6 +----- tests/unit/dialects/test_pg_textsearch.py | 9 ++++----- 5 files changed, 15 insertions(+), 27 deletions(-) diff --git a/sqlspec/dialects/postgres/_generators.py b/sqlspec/dialects/postgres/_generators.py index e74c5dc79..da8454161 100644 --- a/sqlspec/dialects/postgres/_generators.py +++ b/sqlspec/dialects/postgres/_generators.py @@ -24,9 +24,12 @@ __all__ = ("PGTextSearchGenerator", "PGVectorGenerator", "ParadeDBGenerator") _BASE_OPERATOR_TRANSFORM = Postgres.Generator.TRANSFORMS[exp.Operator] -_POSTGRES_EXTENSION_DIALECT_NAMES: Final[frozenset[str]] = frozenset( - {"Postgres", "PGVector", "ParadeDB", "PGTextSearch"} -) +_POSTGRES_EXTENSION_DIALECT_NAMES: Final[frozenset[str]] = frozenset({ + "Postgres", + "PGVector", + "ParadeDB", + "PGTextSearch", +}) def _postgres_extension_operator_sql(generator: PostgresGenerator, expression: exp.Operator) -> str: diff --git a/sqlspec/dialects/postgres/_operators.py b/sqlspec/dialects/postgres/_operators.py index 4ea891a19..780e3fd36 100644 --- a/sqlspec/dialects/postgres/_operators.py +++ b/sqlspec/dialects/postgres/_operators.py @@ -14,8 +14,8 @@ __all__ = ( "PARADEDB_OPERATOR_TOKENS", - "PG_TEXTSEARCH_OPERATOR_TOKENS", "PGVECTOR_OPERATOR_TOKENS", + "PG_TEXTSEARCH_OPERATOR_TOKENS", "is_postgres_extension_operator", "postgres_extension_operator", "register_postgres_extension_operators", @@ -41,9 +41,7 @@ "##": TokenType.NESTED, "##>": TokenType.AGGREGATEFUNCTION, } -PG_TEXTSEARCH_OPERATOR_TOKENS: Final[dict[str, TokenType]] = { - "<@>": TokenType.RING, -} +PG_TEXTSEARCH_OPERATOR_TOKENS: Final[dict[str, TokenType]] = {"<@>": TokenType.RING} _REGISTERED = False @@ -65,11 +63,7 @@ def register_postgres_extension_operators() -> None: return factor: dict[TokenType, Any] = dict(PostgresParser.FACTOR) - extension_tokens = { - **PGVECTOR_OPERATOR_TOKENS, - **PARADEDB_OPERATOR_TOKENS, - **PG_TEXTSEARCH_OPERATOR_TOKENS, - } + extension_tokens = {**PGVECTOR_OPERATOR_TOKENS, **PARADEDB_OPERATOR_TOKENS, **PG_TEXTSEARCH_OPERATOR_TOKENS} for operator, token in extension_tokens.items(): factor[token] = _build_operator_factory(operator) diff --git a/sqlspec/dialects/postgres/_paradedb.py b/sqlspec/dialects/postgres/_paradedb.py index db4411468..1ed67f114 100644 --- a/sqlspec/dialects/postgres/_paradedb.py +++ b/sqlspec/dialects/postgres/_paradedb.py @@ -34,11 +34,7 @@ class ParadeDBTokenizer(Postgres.Tokenizer): """Tokenizer with ParadeDB search operators and pgvector distance operators.""" - KEYWORDS = { - **Postgres.Tokenizer.KEYWORDS, - **PARADEDB_OPERATOR_TOKENS, - **PGVECTOR_OPERATOR_TOKENS, - } + KEYWORDS = {**Postgres.Tokenizer.KEYWORDS, **PARADEDB_OPERATOR_TOKENS, **PGVECTOR_OPERATOR_TOKENS} class ParadeDB(Postgres): diff --git a/sqlspec/dialects/postgres/_pg_textsearch.py b/sqlspec/dialects/postgres/_pg_textsearch.py index 4baf8f8ec..a4afe4610 100644 --- a/sqlspec/dialects/postgres/_pg_textsearch.py +++ b/sqlspec/dialects/postgres/_pg_textsearch.py @@ -28,11 +28,7 @@ class PGTextSearchTokenizer(Postgres.Tokenizer): """Tokenizer with pg_textsearch BM25 ranking operators and pgvector distance operators.""" - KEYWORDS = { - **Postgres.Tokenizer.KEYWORDS, - **PG_TEXTSEARCH_OPERATOR_TOKENS, - **PGVECTOR_OPERATOR_TOKENS, - } + KEYWORDS = {**Postgres.Tokenizer.KEYWORDS, **PG_TEXTSEARCH_OPERATOR_TOKENS, **PGVECTOR_OPERATOR_TOKENS} class PGTextSearch(Postgres): diff --git a/tests/unit/dialects/test_pg_textsearch.py b/tests/unit/dialects/test_pg_textsearch.py index 80b5b20b5..77609b4b8 100644 --- a/tests/unit/dialects/test_pg_textsearch.py +++ b/tests/unit/dialects/test_pg_textsearch.py @@ -7,7 +7,9 @@ def test_pg_textsearch_bm25_ranking_operator() -> None: """Verify BM25 relevance ranking operator <@> parses and generates.""" - sql = "SELECT title, content <@> 'database system' AS score FROM documents ORDER BY content <@> 'database system' ASC" + sql = ( + "SELECT title, content <@> 'database system' AS score FROM documents ORDER BY content <@> 'database system' ASC" + ) expression = parse_one(sql, read=PGTextSearch) rendered = expression.sql(dialect=PGTextSearch) assert "<@>" in rendered @@ -36,10 +38,7 @@ def test_pg_textsearch_hybrid_query() -> None: def test_pg_textsearch_bm25_index_ddl() -> None: """Verify BM25 index creation statement parses.""" - sql = ( - "CREATE INDEX idx_docs_bm25 ON documents " - "USING bm25 (content) WITH (text_config='english', k1=1.2, b=0.75)" - ) + sql = "CREATE INDEX idx_docs_bm25 ON documents USING bm25 (content) WITH (text_config='english', k1=1.2, b=0.75)" expression = parse_one(sql, read=PGTextSearch) rendered = expression.sql(dialect=PGTextSearch) assert "USING bm25" in rendered From ccea283a1332d04cacc2993f4e9485235b422316 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 13 Sep 2026 17:41:00 +0000 Subject: [PATCH 06/10] fix(postgres): preserve extension composition and adapter semantics --- docs/changelog.rst | 25 +++++-------- docs/reference/dialects.rst | 25 ++++++++++++- sqlspec/adapters/adbc/core.py | 15 ++++++-- sqlspec/adapters/adbc/type_converter.py | 12 +++++- sqlspec/core/splitter.py | 1 + sqlspec/dialects/postgres/_operators.py | 15 +++++++- sqlspec/dialects/postgres/_paradedb.py | 10 ++++- .../adk/migrations/0001_create_adk_tables.py | 2 +- .../test_adbc/test_adbc_serialization.py | 7 ++++ tests/unit/adapters/test_adbc/test_core.py | 19 ++++++++++ .../test_adbc/test_extension_detection.py | 19 ++++++++++ tests/unit/core/test_config_runtime.py | 37 +++++++++++++++++++ tests/unit/core/test_splitter.py | 1 + tests/unit/dialects/test_operators.py | 10 ++--- tests/unit/dialects/test_pg_textsearch.py | 17 ++++++++- .../extensions/test_adk/test_migrations.py | 7 ++-- 16 files changed, 185 insertions(+), 37 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index a3d9d6f96..7d267f9ee 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -9,8 +9,8 @@ important operational fixes. Recent Updates ============== -Unreleased ----------- +v0.63.0 - Transactions, table fixtures, SQL fragments, storage, and kwargs parameter binding +--------------------------------------------------------------------------------------------------- **Added:** @@ -22,19 +22,7 @@ Unreleased and adapter modules, with ``active_extensions`` capability tracking on runtime driver features. * Exposed ``pg_textsearch_available`` property across ``AsyncpgConfig``, ``PsycopgSyncConfig``, ``PsycopgAsyncConfig``, ``AdbcConfig``, and ``PsqlpyConfig``. - -**Changed:** - -* Decoupled the ``ParadeDB`` dialect so it inherits directly from ``Postgres`` rather than ``PGVector``, - allowing clean independent combinations of vector search and BM25 extensions. -* Standardized PostgreSQL extension detection across all adapters on a single first-connection probe - via :func:`~sqlspec.core.config_runtime.build_postgres_extension_probe_names`, removing ad-hoc ADK - probe branches. - -v0.63.0 - Transactions, table fixtures, SQL fragments, storage, and kwargs parameter binding ---------------------------------------------------------------------------------------------------- - -**Added:** + (`#782 `_) * Public table-queue primitives extracted to :mod:`sqlspec.extensions.events.primitives` and exported from :mod:`sqlspec.extensions.events`: :func:`~sqlspec.extensions.events.lock_clause`, @@ -165,6 +153,13 @@ v0.63.0 - Transactions, table fixtures, SQL fragments, storage, and kwargs param **Changed:** +* Decoupled the ``ParadeDB`` dialect so it inherits directly from ``Postgres`` rather than ``PGVector``, + allowing clean independent combinations of vector search and BM25 extensions. +* Standardized PostgreSQL extension detection across all adapters on a single first-connection probe + via :func:`~sqlspec.core.config_runtime.build_postgres_extension_probe_names`, removing ad-hoc ADK + probe branches. + (`#782 `_) + * Driver execution methods (:meth:`~sqlspec.driver.SyncDriverAdapterBase.execute`, :meth:`~sqlspec.driver.SyncDriverAdapterBase.select`, etc.) enforce keyword argument parameter passing (``execute(sql, a=1, b=2)`` or ``execute(sql, **params)``). Passing positional dictionary literals diff --git a/docs/reference/dialects.rst b/docs/reference/dialects.rst index 3bdb61aa8..085aa232c 100644 --- a/docs/reference/dialects.rst +++ b/docs/reference/dialects.rst @@ -58,7 +58,10 @@ PGTextSearch :show-inheritance: :no-index: -Adds support for the native PostgreSQL and Google Cloud AlloyDB / AlloyDB Omni ``pg_textsearch`` BM25 extension: +Adds support for PostgreSQL deployments with the ``pg_textsearch`` BM25 extension installed. +AlloyDB currently provides this extension in preview on PostgreSQL 17 and 18; see +`AlloyDB BM25 requirements `_. +Server packaging and version support are independent of SQLSpec dialect availability. .. list-table:: :header-rows: 1 @@ -68,7 +71,25 @@ Adds support for the native PostgreSQL and Google Cloud AlloyDB / AlloyDB Omni ` * - ``<@>`` - BM25 score ranking operator (returns negative score for ASC index scans) -Indices are created with ``USING bm25 (column) WITH (text_config='english')`` and queries order by ``column <@> 'query' ASC``. +Indexes are created with ``USING bm25 (column) WITH (text_config='english')`` and queries order by ``column <@> 'query' ASC``. +Enable the extension in the database before using its operators or index method. +The dialect also supports pgvector distance operators for hybrid queries. + +Asyncpg, psycopg, psqlpy, and PostgreSQL-backed ADBC configurations probe enabled +extensions on first connection. ``enable_pg_textsearch`` defaults to ``True``; +setting it to ``False`` disables detection, not the installed server extension. +``pg_textsearch_available`` and ``is_postgres_extension_active()`` report the cached, +enabled-and-detected state, and remain false before a successful probe. +ADK ``enable_bm25`` requires successful pg_textsearch detection. + +The dialect label remains selected in priority order: ``paradedb``, +``pg_textsearch``, then ``pgvector`` for an otherwise default PostgreSQL +configuration. The active extension set records all enabled discoveries; +ParadeDB includes both text-search operator families so a coinstalled +``pg_search`` does not hide ``pg_textsearch`` syntax. +An explicitly selected non-default dialect is preserved; select a dialect that +supports the operators your queries use. Extension detection does not override +that choice, and a dialect label alone does not mark an extension as available. ParadeDB -------- diff --git a/sqlspec/adapters/adbc/core.py b/sqlspec/adapters/adbc/core.py index 39e2f289f..0d3cad312 100644 --- a/sqlspec/adapters/adbc/core.py +++ b/sqlspec/adapters/adbc/core.py @@ -460,15 +460,21 @@ def resolve_dialect_name(dialect: Any) -> str: """Return the normalized dialect name string.""" if dialect is None: return "" + if isinstance(dialect, str): + return dialect.lower() + if isinstance(dialect, type) and issubclass(dialect, sqlglot.Dialect): + return dialect.__name__.lower() + if isinstance(dialect, sqlglot.Dialect): + return type(dialect).__name__.lower() return str(dialect) def is_postgres_dialect(dialect_name: str) -> bool: """Return True when the dialect indicates PostgreSQL. - Includes pgvector and paradedb which are PostgreSQL extension dialects. + Includes pgvector, paradedb, and pg_textsearch extension dialects. """ - return dialect_name in {"postgres", "postgresql", "pgvector", "paradedb"} + return dialect_name in {"postgres", "postgresql", "pgvector", "paradedb", "pg_textsearch", "pgtextsearch"} def handle_postgres_rollback(dialect: str, cursor: Any, logger: Any | None = None) -> None: @@ -731,7 +737,8 @@ def build_profile() -> "DriverParameterProfile": def get_statement_config(detected_dialect: str) -> StatementConfig: """Create statement configuration for the specified dialect.""" default_style, supported_styles = DIALECT_PARAMETER_STYLES.get( - detected_dialect, (ParameterStyle.QMARK, [ParameterStyle.QMARK]) + "postgres" if is_postgres_dialect(detected_dialect) else detected_dialect, + (ParameterStyle.QMARK, [ParameterStyle.QMARK]), ) sqlglot_dialect = "postgres" if detected_dialect == "postgresql" else detected_dialect @@ -749,7 +756,7 @@ def get_statement_config(detected_dialect: str) -> StatementConfig: parameter_overrides["preserve_parameter_format"] = False parameter_overrides["supported_execution_parameter_styles"] = {ParameterStyle.QMARK, ParameterStyle.NUMERIC} - if detected_dialect in {"postgres", "postgresql"}: + if is_postgres_dialect(detected_dialect): parameter_overrides["ast_transformer"] = build_null_pruning_transform(dialect=sqlglot_dialect) return build_statement_config_from_profile( diff --git a/sqlspec/adapters/adbc/type_converter.py b/sqlspec/adapters/adbc/type_converter.py index 2c96742a8..3612b00da 100644 --- a/sqlspec/adapters/adbc/type_converter.py +++ b/sqlspec/adapters/adbc/type_converter.py @@ -34,7 +34,15 @@ def convert_dict(self, value: "dict[str, Any]") -> Any: Returns: Converted value appropriate for the dialect. """ - if self.dialect in {"postgres", "postgresql", "pgvector", "paradedb", "bigquery"}: + if self.dialect in { + "postgres", + "postgresql", + "pgvector", + "paradedb", + "pg_textsearch", + "pgtextsearch", + "bigquery", + }: return to_json(value) return value @@ -51,7 +59,7 @@ def convert_sequence(self, value: "list[Any] | tuple[Any, ...]") -> "list[Any]": Converted list parameter appropriate for the dialect. """ items = list(value) - if self.dialect in {"postgres", "postgresql", "pgvector", "paradedb"}: + if self.dialect in {"postgres", "postgresql", "pgvector", "paradedb", "pg_textsearch", "pgtextsearch"}: return [item if item is not None else None for item in items] return items diff --git a/sqlspec/core/splitter.py b/sqlspec/core/splitter.py index c84446435..e23592a54 100644 --- a/sqlspec/core/splitter.py +++ b/sqlspec/core/splitter.py @@ -552,6 +552,7 @@ class BigQueryDialectConfig(_EagerDialectConfig): "postgres": PostgreSQLDialectConfig, "paradedb": PostgreSQLDialectConfig, "pg_textsearch": PostgreSQLDialectConfig, + "pgtextsearch": PostgreSQLDialectConfig, "pgvector": PostgreSQLDialectConfig, "mysql": MySQLDialectConfig, "sqlite": SQLiteDialectConfig, diff --git a/sqlspec/dialects/postgres/_operators.py b/sqlspec/dialects/postgres/_operators.py index 780e3fd36..760968773 100644 --- a/sqlspec/dialects/postgres/_operators.py +++ b/sqlspec/dialects/postgres/_operators.py @@ -55,19 +55,30 @@ def _factory(this: exp.Expr | None, expression: exp.Expr | None) -> exp.Operator return _factory +def _parse_pg_textsearch_operator( + _parser: PostgresParser, this: exp.Expr | None, expression: exp.Expr | None +) -> exp.Operator: + node = exp.Operator(this=this, expression=expression, operator="<@>") + node.meta[_CUSTOM_OPERATOR_META_KEY] = "<@>" + return node + + def register_postgres_extension_operators() -> None: - """Patch the compiled Postgres parser with pgvector and ParadeDB operators.""" + """Patch the compiled Postgres parser with PostgreSQL extension operators.""" global _REGISTERED if _REGISTERED: return factor: dict[TokenType, Any] = dict(PostgresParser.FACTOR) - extension_tokens = {**PGVECTOR_OPERATOR_TOKENS, **PARADEDB_OPERATOR_TOKENS, **PG_TEXTSEARCH_OPERATOR_TOKENS} + extension_tokens = {**PGVECTOR_OPERATOR_TOKENS, **PARADEDB_OPERATOR_TOKENS} for operator, token in extension_tokens.items(): factor[token] = _build_operator_factory(operator) setattr(PostgresParser, "FACTOR", factor) + operators = dict(PostgresParser.JSON_OPERATORS) + operators[PG_TEXTSEARCH_OPERATOR_TOKENS["<@>"]] = _parse_pg_textsearch_operator + setattr(PostgresParser, "JSON_OPERATORS", operators) _REGISTERED = True diff --git a/sqlspec/dialects/postgres/_paradedb.py b/sqlspec/dialects/postgres/_paradedb.py index 1ed67f114..5ca022088 100644 --- a/sqlspec/dialects/postgres/_paradedb.py +++ b/sqlspec/dialects/postgres/_paradedb.py @@ -12,7 +12,7 @@ Scoring and snippets are plain functions in ParadeDB (``pdb.score()``, ``pdb.snippet()``), not operators, so they need no dialect support. -Also inherits the pgvector distance operators from PGVector. Registered with +Also supports pgvector distance operators and pg_textsearch ranking independently. Registered with sqlglot through the ``sqlglot.dialects`` entry-point group in ``pyproject.toml`` and by the ``Dialect`` metaclass on import. """ @@ -22,6 +22,7 @@ from sqlspec.dialects.postgres._generators import ParadeDBGenerator from sqlspec.dialects.postgres._operators import ( PARADEDB_OPERATOR_TOKENS, + PG_TEXTSEARCH_OPERATOR_TOKENS, PGVECTOR_OPERATOR_TOKENS, register_postgres_extension_operators, ) @@ -34,7 +35,12 @@ class ParadeDBTokenizer(Postgres.Tokenizer): """Tokenizer with ParadeDB search operators and pgvector distance operators.""" - KEYWORDS = {**Postgres.Tokenizer.KEYWORDS, **PARADEDB_OPERATOR_TOKENS, **PGVECTOR_OPERATOR_TOKENS} + KEYWORDS = { + **Postgres.Tokenizer.KEYWORDS, + **PARADEDB_OPERATOR_TOKENS, + **PGVECTOR_OPERATOR_TOKENS, + **PG_TEXTSEARCH_OPERATOR_TOKENS, + } class ParadeDB(Postgres): diff --git a/sqlspec/extensions/adk/migrations/0001_create_adk_tables.py b/sqlspec/extensions/adk/migrations/0001_create_adk_tables.py index 099cb1eee..5dfd5c46a 100644 --- a/sqlspec/extensions/adk/migrations/0001_create_adk_tables.py +++ b/sqlspec/extensions/adk/migrations/0001_create_adk_tables.py @@ -40,7 +40,7 @@ CREATE_VECTOR_EXTENSION = "CREATE EXTENSION IF NOT EXISTS vector" CREATE_PG_TEXTSEARCH_EXTENSION = "CREATE EXTENSION IF NOT EXISTS pg_textsearch" -_POSTGRES_DIALECTS = frozenset({"postgres", "postgresql", "pgvector", "paradedb"}) +_POSTGRES_DIALECTS = frozenset({"postgres", "postgresql", "pgvector", "paradedb", "pg_textsearch", "pgtextsearch"}) _VECTOR_COLUMN_PATTERN = re.compile(r"\bVECTOR\s*\(", re.IGNORECASE) _BM25_INDEX_PATTERN = re.compile(r"\bUSING\s+bm25\s*\(", re.IGNORECASE) diff --git a/tests/unit/adapters/test_adbc/test_adbc_serialization.py b/tests/unit/adapters/test_adbc/test_adbc_serialization.py index a6e5dc7ca..170ffd1d1 100644 --- a/tests/unit/adapters/test_adbc/test_adbc_serialization.py +++ b/tests/unit/adapters/test_adbc/test_adbc_serialization.py @@ -5,6 +5,7 @@ import pytest from sqlspec.adapters.adbc import AdbcConfig +from sqlspec.adapters.adbc.type_converter import ADBCOutputConverter from sqlspec.utils.serializers import to_json @@ -118,3 +119,9 @@ def test_backward_compatibility_no_serializer() -> None: assert "json_serializer" in config.driver_features assert config.driver_features["json_serializer"] is to_json + + +def test_pg_textsearch_converter_keeps_postgres_json_and_array_handling() -> None: + converter = ADBCOutputConverter("pg_textsearch") + assert converter.convert_dict({"value": 1}) == to_json({"value": 1}) + assert converter.convert_sequence([1, None, 3]) == [1, None, 3] diff --git a/tests/unit/adapters/test_adbc/test_core.py b/tests/unit/adapters/test_adbc/test_core.py index e05f7eca5..e69eeda00 100644 --- a/tests/unit/adapters/test_adbc/test_core.py +++ b/tests/unit/adapters/test_adbc/test_core.py @@ -17,6 +17,7 @@ resolve_column_names, resolve_many_rowcount, ) +from sqlspec.core import ParameterStyle from sqlspec.exceptions import ( DeadlockError, OperationalError, @@ -338,3 +339,21 @@ def test_base_type_coercion_map_replaces_getter_function() -> None: assert isinstance(adbc_core._BASE_TYPE_COERCION_MAP, dict) assert not hasattr(adbc_core, "_get_type_coercion_map") assert adbc_core.driver_profile is not None + + +def test_pg_textsearch_retains_postgres_transaction_and_parameter_behavior() -> None: + executed: list[str] = [] + cursor = SimpleNamespace(execute=lambda statement: executed.append(statement)) + adbc_core.handle_postgres_rollback("pg_textsearch", cursor) + assert executed == ["ROLLBACK"] + assert adbc_core.normalize_postgres_empty_parameters("pg_textsearch", {}) is None + config = get_statement_config("pg_textsearch") + assert config.parameter_config.default_parameter_style == ParameterStyle.NUMERIC + + +def test_pg_textsearch_class_and_instance_resolve_as_postgres_family() -> None: + from sqlspec.dialects import PGTextSearch + + for dialect in (PGTextSearch, PGTextSearch()): + name = adbc_core.resolve_dialect_name(dialect) + assert adbc_core.is_postgres_dialect(name) diff --git a/tests/unit/adapters/test_adbc/test_extension_detection.py b/tests/unit/adapters/test_adbc/test_extension_detection.py index bd39984cd..050acb361 100644 --- a/tests/unit/adapters/test_adbc/test_extension_detection.py +++ b/tests/unit/adapters/test_adbc/test_extension_detection.py @@ -1,5 +1,7 @@ """Unit tests for ADBC postgres extension detection logic.""" +from unittest.mock import MagicMock + from pytest import MonkeyPatch from sqlspec.adapters.adbc.config import AdbcConfig @@ -209,3 +211,20 @@ def fail_create_connection(_self: AdbcConfig) -> None: assert config._pgvector_available is False # pyright: ignore[reportPrivateUsage] assert config._paradedb_available is False # pyright: ignore[reportPrivateUsage] assert config.statement_config.dialect == "sqlite" + + +def test_adbc_explicit_pg_textsearch_dialect_probes_extensions(monkeypatch: MonkeyPatch) -> None: + config = AdbcConfig( + connection_config={"driver_name": "postgres", "uri": "postgresql://localhost/test"}, + statement_config=get_statement_config("pg_textsearch"), + driver_features={"enable_pgvector": False, "enable_paradedb": False}, + ) + connection = MagicMock() + connection.cursor.return_value.fetchall.return_value = [("pg_textsearch",)] + monkeypatch.setattr(AdbcConfig, "create_connection", lambda _self: connection) + + config._detect_extensions_if_needed() # pyright: ignore[reportPrivateUsage] + + assert config.pg_textsearch_available is True + assert config.driver_features["active_extensions"] == {"pg_textsearch"} + connection.close.assert_called_once() diff --git a/tests/unit/core/test_config_runtime.py b/tests/unit/core/test_config_runtime.py index 34c66f4e5..232f7438f 100644 --- a/tests/unit/core/test_config_runtime.py +++ b/tests/unit/core/test_config_runtime.py @@ -1,5 +1,11 @@ """Unit tests for sqlspec.core.config_runtime.""" +from typing import Any + +import pytest +from sqlglot import parse_one + +import sqlspec.dialects.postgres # noqa: F401 from sqlspec.core.config_runtime import ( build_postgres_extension_probe_names, is_postgres_extension_active, @@ -29,3 +35,34 @@ def test_resolve_postgres_extension_state_active_extensions() -> None: assert is_postgres_extension_active(features, "pg_textsearch") is True assert is_postgres_extension_active(features, "vector") is True assert is_postgres_extension_active(features, "pg_search") is False + + +@pytest.mark.parametrize( + "detected", + [ + set(), + {"vector"}, + {"pg_search"}, + {"pg_textsearch"}, + {"vector", "pg_search"}, + {"vector", "pg_textsearch"}, + {"pg_search", "pg_textsearch"}, + {"vector", "pg_search", "pg_textsearch"}, + ], +) +def test_enabled_extension_combinations_preserve_all_operators(detected: set[str]) -> None: + features: dict[str, Any] = {"enable_pgvector": True, "enable_paradedb": True, "enable_pg_textsearch": True} + config, _, _ = resolve_postgres_extension_state(StatementConfig(dialect="postgres"), features, detected) + expressions = ["1"] + if "vector" in detected: + expressions.append("embedding <=> '[1,2,3]'") + if "pg_search" in detected: + expressions.append("content @@@ 'query'") + if "pg_textsearch" in detected: + expressions.append("content <@> 'query'") + query = "SELECT " + ", ".join(expressions) + " FROM documents" + rendered = parse_one(query, read=config.dialect).sql(dialect=config.dialect) + for operator in ("<=>", "@@@", "<@>"): + if operator in query: + assert operator in rendered + assert features["active_extensions"] == detected diff --git a/tests/unit/core/test_splitter.py b/tests/unit/core/test_splitter.py index 99b39d64e..fa2fbd025 100644 --- a/tests/unit/core/test_splitter.py +++ b/tests/unit/core/test_splitter.py @@ -46,6 +46,7 @@ def test_dialect_class_map_contains_expected_aliases() -> None: "paradedb": splitter_module.PostgreSQLDialectConfig, "pgvector": splitter_module.PostgreSQLDialectConfig, "pg_textsearch": splitter_module.PostgreSQLDialectConfig, + "pgtextsearch": splitter_module.PostgreSQLDialectConfig, "mysql": splitter_module.MySQLDialectConfig, "sqlite": splitter_module.SQLiteDialectConfig, "duckdb": splitter_module.DuckDBDialectConfig, diff --git a/tests/unit/dialects/test_operators.py b/tests/unit/dialects/test_operators.py index 89619568c..c7f9ff36f 100644 --- a/tests/unit/dialects/test_operators.py +++ b/tests/unit/dialects/test_operators.py @@ -20,16 +20,16 @@ def test_pg_textsearch_operator_tokens_definition() -> None: assert PG_TEXTSEARCH_OPERATOR_TOKENS["<@>"] == TokenType.RING -def test_pg_textsearch_factor_registration() -> None: - """Verify pg_textsearch operator is registered in PostgresParser.FACTOR.""" +def test_pg_textsearch_operator_registration() -> None: + """Verify pg_textsearch operator is registered in PostgresParser.JSON_OPERATORS.""" register_postgres_extension_operators() ring_token = PG_TEXTSEARCH_OPERATOR_TOKENS["<@>"] - assert ring_token in PostgresParser.FACTOR + assert ring_token in PostgresParser.JSON_OPERATORS - factory = PostgresParser.FACTOR[ring_token] + factory = PostgresParser.JSON_OPERATORS[ring_token] left = exp.var("content") right = exp.Literal.string("search query") - node = factory(left, right) + node = factory(None, left, right) assert isinstance(node, exp.Operator) assert is_postgres_extension_operator(node) diff --git a/tests/unit/dialects/test_pg_textsearch.py b/tests/unit/dialects/test_pg_textsearch.py index 77609b4b8..88886ca57 100644 --- a/tests/unit/dialects/test_pg_textsearch.py +++ b/tests/unit/dialects/test_pg_textsearch.py @@ -1,6 +1,6 @@ """Unit tests for the PGTextSearch sqlglot dialect.""" -from sqlglot import parse_one +from sqlglot import exp, parse_one from sqlspec.dialects.postgres import PGTextSearch @@ -43,3 +43,18 @@ def test_pg_textsearch_bm25_index_ddl() -> None: rendered = expression.sql(dialect=PGTextSearch) assert "USING bm25" in rendered assert "text_config" in rendered + + +def test_pg_textsearch_uses_postgres_custom_operator_precedence() -> None: + expression = parse_one("SELECT title || content <@> 'query' FROM documents", read=PGTextSearch) + ranking = expression.expressions[0] + assert isinstance(ranking, exp.Operator) + assert isinstance(ranking.this, exp.DPipe) + assert ranking.expression == exp.Literal.string("query") + + +def test_pg_textsearch_preserves_postgres_containment_operators() -> None: + expression = parse_one("SELECT ARRAY[1] <@ ARRAY[1, 2], ARRAY[1, 2] @> ARRAY[1]", read=PGTextSearch) + assert "<@" in expression.sql(dialect=PGTextSearch) + assert "@>" in expression.sql(dialect=PGTextSearch) + assert not list(expression.find_all(exp.Operator)) diff --git a/tests/unit/extensions/test_adk/test_migrations.py b/tests/unit/extensions/test_adk/test_migrations.py index 4716e88bb..c788e2b0f 100644 --- a/tests/unit/extensions/test_adk/test_migrations.py +++ b/tests/unit/extensions/test_adk/test_migrations.py @@ -102,7 +102,7 @@ async def test_create_migration_resolves_only_the_enabled_store_class(monkeypatc assert memory_calls == ["_get_memory_store_class"] -@pytest.mark.parametrize("dialect", ["postgres", "postgresql", "pgvector", "paradedb"]) +@pytest.mark.parametrize("dialect", ["postgres", "postgresql", "pgvector", "paradedb", "pg_textsearch", "pgtextsearch"]) async def test_asyncpg_memory_migration_installs_vector_extension_before_first_use(dialect: str) -> None: """PostgreSQL vector memory DDL is preceded by exactly one extension statement.""" context = MigrationContext(config=_asyncpg_config(), dialect=dialect) @@ -130,9 +130,10 @@ async def test_psycopg_memory_migration_installs_vector_extension() -> None: ) -async def test_postgres_bm25_migration_installs_pg_textsearch_extension() -> None: +@pytest.mark.parametrize("dialect", ["postgres", "pg_textsearch", "pgtextsearch"]) +async def test_postgres_bm25_migration_installs_pg_textsearch_extension(dialect: str) -> None: """BM25 DDL is preceded by an idempotent pg_textsearch enablement statement.""" - context = MigrationContext(config=_asyncpg_config({"enable_bm25": True}), dialect="postgres") + context = MigrationContext(config=_asyncpg_config({"enable_bm25": True}), dialect=dialect) statements = await migration.up(context) From 923e882ab4e467bcaa30639bfe85e19b7bc7e1b6 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 13 Sep 2026 18:29:09 +0000 Subject: [PATCH 07/10] fix(dialects): decouple paradedb operators and align dialect tests with prd --- docs/reference/dialects.rst | 4 +- sqlspec/dialects/postgres/_paradedb.py | 10 +--- tests/unit/core/test_config_runtime.py | 68 +++++++++++++++++++------- 3 files changed, 52 insertions(+), 30 deletions(-) diff --git a/docs/reference/dialects.rst b/docs/reference/dialects.rst index 085aa232c..1ac982d50 100644 --- a/docs/reference/dialects.rst +++ b/docs/reference/dialects.rst @@ -84,9 +84,7 @@ ADK ``enable_bm25`` requires successful pg_textsearch detection. The dialect label remains selected in priority order: ``paradedb``, ``pg_textsearch``, then ``pgvector`` for an otherwise default PostgreSQL -configuration. The active extension set records all enabled discoveries; -ParadeDB includes both text-search operator families so a coinstalled -``pg_search`` does not hide ``pg_textsearch`` syntax. +configuration. The active extension set records all enabled discoveries independently. An explicitly selected non-default dialect is preserved; select a dialect that supports the operators your queries use. Extension detection does not override that choice, and a dialect label alone does not mark an extension as available. diff --git a/sqlspec/dialects/postgres/_paradedb.py b/sqlspec/dialects/postgres/_paradedb.py index 5ca022088..a403ca106 100644 --- a/sqlspec/dialects/postgres/_paradedb.py +++ b/sqlspec/dialects/postgres/_paradedb.py @@ -12,7 +12,7 @@ Scoring and snippets are plain functions in ParadeDB (``pdb.score()``, ``pdb.snippet()``), not operators, so they need no dialect support. -Also supports pgvector distance operators and pg_textsearch ranking independently. Registered with +Also supports pgvector distance operators independently. Registered with sqlglot through the ``sqlglot.dialects`` entry-point group in ``pyproject.toml`` and by the ``Dialect`` metaclass on import. """ @@ -22,7 +22,6 @@ from sqlspec.dialects.postgres._generators import ParadeDBGenerator from sqlspec.dialects.postgres._operators import ( PARADEDB_OPERATOR_TOKENS, - PG_TEXTSEARCH_OPERATOR_TOKENS, PGVECTOR_OPERATOR_TOKENS, register_postgres_extension_operators, ) @@ -35,12 +34,7 @@ class ParadeDBTokenizer(Postgres.Tokenizer): """Tokenizer with ParadeDB search operators and pgvector distance operators.""" - KEYWORDS = { - **Postgres.Tokenizer.KEYWORDS, - **PARADEDB_OPERATOR_TOKENS, - **PGVECTOR_OPERATOR_TOKENS, - **PG_TEXTSEARCH_OPERATOR_TOKENS, - } + KEYWORDS = {**Postgres.Tokenizer.KEYWORDS, **PARADEDB_OPERATOR_TOKENS, **PGVECTOR_OPERATOR_TOKENS} class ParadeDB(Postgres): diff --git a/tests/unit/core/test_config_runtime.py b/tests/unit/core/test_config_runtime.py index 232f7438f..f86d8d409 100644 --- a/tests/unit/core/test_config_runtime.py +++ b/tests/unit/core/test_config_runtime.py @@ -21,7 +21,7 @@ def test_build_postgres_extension_probe_names_pg_textsearch() -> None: def test_resolve_postgres_extension_state_active_extensions() -> None: - features = {"enable_pgvector": True, "enable_pg_textsearch": True} + features: dict[str, Any] = {"enable_pgvector": True, "enable_pg_textsearch": True} config = StatementConfig(dialect="postgres") detected = {"vector", "pg_textsearch"} @@ -38,31 +38,61 @@ def test_resolve_postgres_extension_state_active_extensions() -> None: @pytest.mark.parametrize( - "detected", + ("detected", "expected_dialect"), [ - set(), - {"vector"}, - {"pg_search"}, - {"pg_textsearch"}, - {"vector", "pg_search"}, - {"vector", "pg_textsearch"}, - {"pg_search", "pg_textsearch"}, - {"vector", "pg_search", "pg_textsearch"}, + (set(), "postgres"), + ({"vector"}, "pgvector"), + ({"pg_search"}, "paradedb"), + ({"pg_textsearch"}, "pg_textsearch"), + ({"vector", "pg_search"}, "paradedb"), + ({"vector", "pg_textsearch"}, "pg_textsearch"), + ({"pg_search", "pg_textsearch"}, "paradedb"), + ({"vector", "pg_search", "pg_textsearch"}, "paradedb"), ], ) -def test_enabled_extension_combinations_preserve_all_operators(detected: set[str]) -> None: +def test_enabled_extension_combinations_promoted_dialect_and_active_extensions( + detected: set[str], expected_dialect: str +) -> None: + """Verify dialect promotion hierarchy and active extension recording for all combinations.""" features: dict[str, Any] = {"enable_pgvector": True, "enable_paradedb": True, "enable_pg_textsearch": True} - config, _, _ = resolve_postgres_extension_state(StatementConfig(dialect="postgres"), features, detected) + config, pgv_avail, pdb_avail = resolve_postgres_extension_state( + StatementConfig(dialect="postgres"), features, detected + ) + assert config.dialect == expected_dialect + assert features["active_extensions"] == detected + assert pgv_avail is ("vector" in detected) + assert pdb_avail is ("pg_search" in detected) + assert is_postgres_extension_active(features, "pg_textsearch") is ("pg_textsearch" in detected) + assert is_postgres_extension_active(features, "vector") is ("vector" in detected) + assert is_postgres_extension_active(features, "pg_search") is ("pg_search" in detected) + expressions = ["1"] - if "vector" in detected: - expressions.append("embedding <=> '[1,2,3]'") - if "pg_search" in detected: + if config.dialect == "paradedb": expressions.append("content @@@ 'query'") - if "pg_textsearch" in detected: + expressions.append("embedding <=> '[1,2,3]'") + elif config.dialect == "pg_textsearch": expressions.append("content <@> 'query'") + expressions.append("embedding <=> '[1,2,3]'") + elif config.dialect == "pgvector": + expressions.append("embedding <=> '[1,2,3]'") + query = "SELECT " + ", ".join(expressions) + " FROM documents" rendered = parse_one(query, read=config.dialect).sql(dialect=config.dialect) - for operator in ("<=>", "@@@", "<@>"): - if operator in query: - assert operator in rendered + for part in expressions[1:]: + op = part.split()[1] + assert op in rendered + + +def test_resolve_postgres_extension_state_preserves_explicit_dialect() -> None: + """Verify an explicit non-postgres dialect is preserved even when other extensions are detected.""" + features: dict[str, Any] = {"enable_pgvector": True, "enable_paradedb": True, "enable_pg_textsearch": True} + config = StatementConfig(dialect="pg_textsearch") + detected = {"vector", "pg_search", "pg_textsearch"} + + updated_config, pgv_avail, pdb_avail = resolve_postgres_extension_state(config, features, detected) + + assert updated_config.dialect == "pg_textsearch" + assert pgv_avail is True + assert pdb_avail is True assert features["active_extensions"] == detected + assert is_postgres_extension_active(features, "pg_textsearch") is True From e4c367aaa1253a325d9f13cfad3d13bac6008925 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 13 Sep 2026 18:41:06 +0000 Subject: [PATCH 08/10] fix(mssql): render standard TIMESTAMP as DATETIME2 on T-SQL DDL --- sqlspec/builder/_ddl.py | 6 +++++- tests/unit/builder/test_ddl_builder.py | 8 ++++++++ 2 files changed, 13 insertions(+), 1 deletion(-) diff --git a/sqlspec/builder/_ddl.py b/sqlspec/builder/_ddl.py index 9742e078f..99293bc1c 100644 --- a/sqlspec/builder/_ddl.py +++ b/sqlspec/builder/_ddl.py @@ -77,8 +77,12 @@ def _parse_column_type(name: str | None, dtype: str, dialect: "DialectType | None") -> exp.DataType: + target_dialect = dialect + norm_dialect = _normalize_dialect(dialect) if dialect else None + if norm_dialect in ("tsql", "mssql") and dtype.strip().upper().startswith("TIMESTAMP"): + target_dialect = None try: - return exp.DataType.build(dtype, dialect=dialect) + return exp.DataType.build(dtype, dialect=target_dialect) except ParseError as exc: msg = f"Column {name!r}: cannot parse type {dtype!r} for dialect {dialect!r}" raise SQLBuilderError(msg) from exc diff --git a/tests/unit/builder/test_ddl_builder.py b/tests/unit/builder/test_ddl_builder.py index 4a476153a..06be3aa75 100644 --- a/tests/unit/builder/test_ddl_builder.py +++ b/tests/unit/builder/test_ddl_builder.py @@ -209,3 +209,11 @@ def test_ddl_accepts_public_dialect_aliases(dialect: str, dtype: str) -> None: builder = sql.create_table("t").column("a", dtype) assert "CREATE TABLE" in builder.build(dialect=dialect).sql assert "CREATE TABLE" in builder.to_statement(StatementConfig(dialect=dialect)).sql + + +def test_create_table_timestamp_on_tsql_renders_datetime2() -> None: + """Standard TIMESTAMP logical column type renders as DATETIME2 on T-SQL instead of ROWVERSION.""" + builder = sql.create_table("t").column("applied_at", "TIMESTAMP", default="CURRENT_TIMESTAMP") + sql_text = builder.build(dialect="tsql").sql + assert "DATETIME2" in sql_text + assert "ROWVERSION" not in sql_text From d3e4085b59034ffa69292660d0d8fd430e819452 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 13 Sep 2026 19:19:06 +0000 Subject: [PATCH 09/10] test(postgres): register pg_textsearch driver feature consumption --- tests/integration/adapters/_shared/_driver_type_system.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/integration/adapters/_shared/_driver_type_system.py b/tests/integration/adapters/_shared/_driver_type_system.py index 85ec8a7f9..6ff021238 100644 --- a/tests/integration/adapters/_shared/_driver_type_system.py +++ b/tests/integration/adapters/_shared/_driver_type_system.py @@ -176,6 +176,7 @@ class SourceEquivalenceCase: "enable_cast_detection", "enable_pgvector", "enable_paradedb", + "enable_pg_textsearch", "json_serializer", "json_deserializer", "on_connection_create", @@ -185,6 +186,7 @@ class SourceEquivalenceCase: "psycopg": ( "enable_pgvector", "enable_paradedb", + "enable_pg_textsearch", "json_serializer", "json_deserializer", "on_connection_create", From 5192045ebdfd59bcc9d7515c3208b8498c495fe1 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 13 Sep 2026 19:32:02 +0000 Subject: [PATCH 10/10] ci: allow cold setup before the integration watchdog --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index a7e6d9fd5..9b47f6b34 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -254,7 +254,7 @@ jobs: fail-fast: true matrix: python-version: ["3.10", "3.11", "3.12", "3.13", "3.14"] - timeout-minutes: 30 + timeout-minutes: 45 steps: - name: Check out repository uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1