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 diff --git a/docs/changelog.rst b/docs/changelog.rst index 078016b28..7d267f9ee 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -14,6 +14,16 @@ v0.63.0 - Transactions, table fixtures, SQL fragments, storage, and kwargs param **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``. + (`#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`, :func:`~sqlspec.extensions.events.row_limit_clause`, @@ -143,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/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..1ac982d50 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,45 @@ Adds support for pgvector distance operators: * - ``<%>`` - Jaccard distance (binary vectors) +PGTextSearch +------------ + +.. autoclass:: sqlspec.dialects.postgres.PGTextSearch + :members: + :show-inheritance: + :no-index: + +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 + + * - Operator + - Description + * - ``<@>`` + - BM25 score ranking operator (returns negative score for ASC index scans) + +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 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. + ParadeDB -------- @@ -58,7 +97,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 5ff276064..bb1aad4cc 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" @@ -236,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/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..0d3cad312 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: @@ -455,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: @@ -726,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 @@ -744,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( @@ -775,6 +787,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/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/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/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/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/sqlspec/core/splitter.py b/sqlspec/core/splitter.py index e27868f6f..e23592a54 100644 --- a/sqlspec/core/splitter.py +++ b/sqlspec/core/splitter.py @@ -551,6 +551,8 @@ class BigQueryDialectConfig(_EagerDialectConfig): "postgresql": PostgreSQLDialectConfig, "postgres": PostgreSQLDialectConfig, "paradedb": PostgreSQLDialectConfig, + "pg_textsearch": PostgreSQLDialectConfig, + "pgtextsearch": 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..da8454161 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,22 @@ ) 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 +61,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..760968773 100644 --- a/sqlspec/dialects/postgres/_operators.py +++ b/sqlspec/dialects/postgres/_operators.py @@ -15,6 +15,7 @@ __all__ = ( "PARADEDB_OPERATOR_TOKENS", "PGVECTOR_OPERATOR_TOKENS", + "PG_TEXTSEARCH_OPERATOR_TOKENS", "is_postgres_extension_operator", "postgres_extension_operator", "register_postgres_extension_operators", @@ -40,6 +41,7 @@ "##": TokenType.NESTED, "##>": TokenType.AGGREGATEFUNCTION, } +PG_TEXTSEARCH_OPERATOR_TOKENS: Final[dict[str, TokenType]] = {"<@>": TokenType.RING} _REGISTERED = False @@ -53,18 +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) - for operator, token in {**PGVECTOR_OPERATOR_TOKENS, **PARADEDB_OPERATOR_TOKENS}.items(): + 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 d6c813b00..a403ca106 100644 --- a/sqlspec/dialects/postgres/_paradedb.py +++ b/sqlspec/dialects/postgres/_paradedb.py @@ -12,27 +12,32 @@ 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 independently. 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 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..a4afe4610 --- /dev/null +++ b/sqlspec/dialects/postgres/_pg_textsearch.py @@ -0,0 +1,38 @@ +"""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/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/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", 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 921e4f036..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 @@ -70,51 +72,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 +149,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" @@ -182,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/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.""" 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 diff --git a/tests/unit/core/test_config_runtime.py b/tests/unit/core/test_config_runtime.py new file mode 100644 index 000000000..f86d8d409 --- /dev/null +++ b/tests/unit/core/test_config_runtime.py @@ -0,0 +1,98 @@ +"""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, + 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: dict[str, Any] = {"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 + + +@pytest.mark.parametrize( + ("detected", "expected_dialect"), + [ + (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_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, 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 config.dialect == "paradedb": + expressions.append("content @@@ 'query'") + 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 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 diff --git a/tests/unit/core/test_splitter.py b/tests/unit/core/test_splitter.py index c60c4a2b1..fa2fbd025 100644 --- a/tests/unit/core/test_splitter.py +++ b/tests/unit/core/test_splitter.py @@ -45,6 +45,8 @@ 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, + "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 new file mode 100644 index 000000000..c7f9ff36f --- /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_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.JSON_OPERATORS + + factory = PostgresParser.JSON_OPERATORS[ring_token] + left = exp.var("content") + right = exp.Literal.string("search query") + node = factory(None, 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..88886ca57 --- /dev/null +++ b/tests/unit/dialects/test_pg_textsearch.py @@ -0,0 +1,60 @@ +"""Unit tests for the PGTextSearch sqlglot dialect.""" + +from sqlglot import exp, 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 + + +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/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) 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) 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",