From d26d1f575708050e2c10de1f260d403d067fd55e Mon Sep 17 00:00:00 2001 From: Erik Johnston Date: Thu, 2 Jul 2026 16:49:52 +0000 Subject: [PATCH] Make PostgresEngine a driver-agnostic base for psycopg2 and Rust Previously RustPostgresEngine subclassed the psycopg2 PostgresEngine, which dragged psycopg2 into the Rust path and conflated "is Postgres" with "is psycopg2". Split the hierarchy so both drivers are siblings: - `PostgresEngine` (new `postgres_base` module) is now the driver-agnostic base holding the shared SQL-dialect and config logic (single_threaded, supports_using_any_list, row_id_name, get_db_locale, check_new_database, lock_table, synchronous_commit / statement_timeout). It has no driver dependency, so it always imports. Everything touching a live connection, the DBAPI2 exception module, or the placeholder style is left abstract. - `Psycopg2Engine(PostgresEngine)` holds the psycopg2 specifics (register_type / register_adapter, isolation-level map, conn.status/closed/server_version, `%s` placeholders, psycopg2 execute path, `uses_psycopg2_extras = True`). - `RustPostgresEngine(PostgresEngine)` is re-parented onto the base (no longer inherits psycopg2). It passes the Rust DBAPI2 module to the base and gets NotImplementedError stubs for the still-psycopg2-shaped check_database / server_version (part of the deferred startup wiring). The base keeps the name `PostgresEngine`, so all ~91 `isinstance(engine, PostgresEngine)` checks across the storage layer (which mean "emit Postgres SQL") hold for both drivers unchanged. `create_engine` now returns `Psycopg2Engine` for `name == "psycopg2"`. Verified: full lint clean; Rust engine/adapter tests pass; a psycopg2 homeserver boots and runs (test_room_search under Postgres) and the sqlite path is unaffected (test_room_search / test_event_federation under sqlite). Co-Authored-By: Claude Opus 4.8 (1M context) Claude-Session: https://claude.ai/code/session_01W3G4M92AmwSSZCbmtMJU3d --- synapse/storage/engines/__init__.py | 21 ++-- synapse/storage/engines/postgres.py | 74 ++------------ synapse/storage/engines/postgres_base.py | 117 +++++++++++++++++++++++ synapse/storage/engines/postgres_rust.py | 35 +++++-- tests/storage/test_room_search.py | 4 +- 5 files changed, 166 insertions(+), 85 deletions(-) create mode 100644 synapse/storage/engines/postgres_base.py diff --git a/synapse/storage/engines/__init__.py b/synapse/storage/engines/__init__.py index 53b4cf3d13..b66afaf913 100644 --- a/synapse/storage/engines/__init__.py +++ b/synapse/storage/engines/__init__.py @@ -22,16 +22,20 @@ from typing import Any, Mapping, NoReturn from ._base import BaseDatabaseEngine, IncorrectDatabaseSetup -# The classes `PostgresEngine` and `Sqlite3Engine` must always be importable, because -# we use `isinstance(engine, PostgresEngine)` to write different queries for postgres -# and sqlite. But the database driver modules are both optional: they may not be -# installed. To account for this, create dummy classes on import failure so we can -# still run `isinstance()` checks. +# `PostgresEngine` is the driver-agnostic Postgres base (psycopg2 and the native +# Rust backend both subclass it). It has no driver dependency, so it always +# imports — which matters because `isinstance(engine, PostgresEngine)` is used +# throughout the storage layer to write Postgres- vs sqlite-flavoured queries. +from .postgres_base import PostgresEngine + +# The concrete driver engines are optional: their driver modules may not be +# installed. Create dummy classes on import failure so `isinstance()` checks +# still work (and construction fails with a clear message). try: - from .postgres import PostgresEngine + from .postgres import Psycopg2Engine except ImportError: - class PostgresEngine(BaseDatabaseEngine): # type: ignore[no-redef] + class Psycopg2Engine(PostgresEngine): # type: ignore[no-redef] def __new__(cls, *args: object, **kwargs: object) -> NoReturn: raise RuntimeError( f"Cannot create {cls.__name__} -- psycopg2 module is not installed" @@ -56,7 +60,7 @@ def create_engine(database_config: Mapping[str, Any]) -> BaseDatabaseEngine: return Sqlite3Engine(database_config) if name == "psycopg2": - return PostgresEngine(database_config) + return Psycopg2Engine(database_config) raise RuntimeError("Unsupported database engine '%s'" % (name,)) @@ -65,6 +69,7 @@ __all__ = [ "create_engine", "BaseDatabaseEngine", "PostgresEngine", + "Psycopg2Engine", "Sqlite3Engine", "IncorrectDatabaseSetup", ] diff --git a/synapse/storage/engines/postgres.py b/synapse/storage/engines/postgres.py index 7cd50fb8f1..972e2cb471 100644 --- a/synapse/storage/engines/postgres.py +++ b/synapse/storage/engines/postgres.py @@ -20,18 +20,16 @@ # import logging -from typing import TYPE_CHECKING, Any, Mapping, NoReturn, cast +from typing import TYPE_CHECKING, Any, Mapping, NoReturn import psycopg2.extensions from synapse.storage.engines._base import ( AUTO_INCREMENT_PRIMARY_KEYPLACEHOLDER, - BaseDatabaseEngine, IncorrectDatabaseSetup, IsolationLevel, ) -from synapse.storage.types import Cursor -from synapse.util.duration import Duration +from synapse.storage.engines.postgres_base import PostgresEngine if TYPE_CHECKING: from synapse.storage.database import LoggingDatabaseConnection @@ -40,9 +38,11 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) -class PostgresEngine( - BaseDatabaseEngine[psycopg2.extensions.connection, psycopg2.extensions.cursor] +class Psycopg2Engine( + PostgresEngine[psycopg2.extensions.connection, psycopg2.extensions.cursor] ): + """The Postgres backend that talks to the database via psycopg2.""" + def __init__(self, database_config: Mapping[str, Any]): super().__init__(psycopg2, database_config) psycopg2.extensions.register_type(psycopg2.extensions.UNICODE) @@ -54,17 +54,6 @@ class PostgresEngine( raise Exception("Passing bytes to DB is disabled.") psycopg2.extensions.register_adapter(bytes, _disable_bytes_adapter) - self.synchronous_commit: bool = database_config.get("synchronous_commit", True) - # Set the statement timeout to 10 minutes by default. - # - # Any query taking more than 10 minutes should probably be considered a bug; - # most of the time this is a sign that work needs to be split up or that - # some degenerate query plan has been created and the client has probably - # timed out/walked off anyway. - # This is in milliseconds. - self.statement_timeout: int | None = database_config.get( - "statement_timeout", Duration(minutes=10).as_millis() - ) self._version: int | None = None # unknown as yet self.isolation_level_map: Mapping[int, int] = { @@ -75,18 +64,6 @@ class PostgresEngine( self.default_isolation_level = ( psycopg2.extensions.ISOLATION_LEVEL_REPEATABLE_READ ) - self.config = database_config - - @property - def single_threaded(self) -> bool: - return False - - def get_db_locale(self, txn: Cursor) -> tuple[str, str]: - txn.execute( - "SELECT datcollate, datctype FROM pg_database WHERE datname = current_database()" - ) - collation, ctype = cast(tuple[str, str], txn.fetchone()) - return collation, ctype def check_database( self, @@ -140,33 +117,6 @@ class PostgresEngine( ctype, ) - def check_new_database(self, txn: Cursor) -> None: - """Gets called when setting up a brand new database. This allows us to - apply stricter checks on new databases versus existing database. - """ - - allow_unsafe_locale = self.config.get("allow_unsafe_locale", False) - if allow_unsafe_locale: - return - - collation, ctype = self.get_db_locale(txn) - - errors = [] - - if collation != "C": - errors.append(" - 'COLLATE' is set to %r. Should be 'C'" % (collation,)) - - if ctype != "C": - errors.append(" - 'CTYPE' is set to %r. Should be 'C'" % (ctype,)) - - if errors: - raise IncorrectDatabaseSetup( - "Database is incorrectly configured:\n\n%s\n\n" - "See docs/postgres.md for more information. You can override this check by" - "setting 'allow_unsafe_locale' to true in the database config.", - "\n".join(errors), - ) - def convert_param_style(self, sql: str) -> str: return sql.replace("?", "%s") @@ -190,11 +140,6 @@ class PostgresEngine( cursor.close() db_conn.commit() - @property - def supports_using_any_list(self) -> bool: - """Do we support using `a = ANY(?)` and passing a list""" - return True - def is_deadlock(self, error: Exception) -> bool: if isinstance(error, psycopg2.DatabaseError): # https://www.postgresql.org/docs/current/static/errcodes-appendix.html @@ -206,9 +151,6 @@ class PostgresEngine( def is_connection_closed(self, conn: psycopg2.extensions.connection) -> bool: return bool(conn.closed) - def lock_table(self, txn: Cursor, table: str) -> None: - txn.execute("LOCK TABLE %s in EXCLUSIVE MODE" % (table,)) - @property def server_version(self) -> str: """Returns a string giving the server version. For example: '8.1.5'.""" @@ -223,10 +165,6 @@ class PostgresEngine( else: return "%i.%i.%i" % (numver / 10000, (numver % 10000) / 100, numver % 100) - @property - def row_id_name(self) -> str: - return "ctid" - def in_transaction(self, conn: psycopg2.extensions.connection) -> bool: return conn.status != psycopg2.extensions.STATUS_READY diff --git a/synapse/storage/engines/postgres_base.py b/synapse/storage/engines/postgres_base.py new file mode 100644 index 0000000000..9a5369352c --- /dev/null +++ b/synapse/storage/engines/postgres_base.py @@ -0,0 +1,117 @@ +# +# This file is licensed under the Affero General Public License (AGPL) version 3. +# +# Copyright 2015, 2016 OpenMarket Ltd +# Copyright (C) 2023 New Vector, Ltd +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU Affero General Public License as +# published by the Free Software Foundation, either version 3 of the +# License, or (at your option) any later version. +# +# See the GNU Affero General Public License for more details: +# . +# +# Originally licensed under the Apache License, Version 2.0: +# . +# +# [This file includes modifications made by New Vector Limited] +# +# + +import logging +from typing import Any, Mapping, TypeVar, cast + +from synapse.storage.engines._base import BaseDatabaseEngine, IncorrectDatabaseSetup +from synapse.storage.types import Connection, Cursor, DBAPI2Module +from synapse.util.duration import Duration + +logger = logging.getLogger(__name__) + +ConnectionType = TypeVar("ConnectionType", bound=Connection) +CursorType = TypeVar("CursorType", bound=Cursor) + + +class PostgresEngine(BaseDatabaseEngine[ConnectionType, CursorType]): + """Behaviour shared by the Postgres backends, regardless of driver. + + This holds the SQL-dialect and configuration logic that is identical whether + Synapse talks to Postgres via psycopg2 (:class:`Psycopg2Engine`) or the + native Rust driver (:class:`RustPostgresEngine`). Everything that touches a + live connection, the DBAPI2 exception module, or the parameter-placeholder + style is left abstract for those subclasses to provide. + + Crucially the name is kept as ``PostgresEngine`` so the many + ``isinstance(engine, PostgresEngine)`` checks across the storage layer — all + of which mean "emit Postgres SQL" — hold for both drivers. + """ + + # Whether `LoggingTransaction` may use the `psycopg2.extras` helpers on the + # cursor. Set by each concrete subclass (True for psycopg2, False for Rust). + uses_psycopg2_extras: bool + + def __init__(self, module: DBAPI2Module, database_config: Mapping[str, Any]): + super().__init__(module, database_config) + + self.synchronous_commit: bool = database_config.get("synchronous_commit", True) + # Set the statement timeout to 10 minutes by default. + # + # Any query taking more than 10 minutes should probably be considered a bug; + # most of the time this is a sign that work needs to be split up or that + # some degenerate query plan has been created and the client has probably + # timed out/walked off anyway. + # This is in milliseconds. + self.statement_timeout: int | None = database_config.get( + "statement_timeout", Duration(minutes=10).as_millis() + ) + self.config = database_config + + @property + def single_threaded(self) -> bool: + return False + + @property + def supports_using_any_list(self) -> bool: + """Do we support using `a = ANY(?)` and passing a list""" + return True + + @property + def row_id_name(self) -> str: + return "ctid" + + def get_db_locale(self, txn: Cursor) -> tuple[str, str]: + txn.execute( + "SELECT datcollate, datctype FROM pg_database WHERE datname = current_database()" + ) + collation, ctype = cast(tuple[str, str], txn.fetchone()) + return collation, ctype + + def check_new_database(self, txn: Cursor) -> None: + """Gets called when setting up a brand new database. This allows us to + apply stricter checks on new databases versus existing database. + """ + + allow_unsafe_locale = self.config.get("allow_unsafe_locale", False) + if allow_unsafe_locale: + return + + collation, ctype = self.get_db_locale(txn) + + errors = [] + + if collation != "C": + errors.append(" - 'COLLATE' is set to %r. Should be 'C'" % (collation,)) + + if ctype != "C": + errors.append(" - 'CTYPE' is set to %r. Should be 'C'" % (ctype,)) + + if errors: + raise IncorrectDatabaseSetup( + "Database is incorrectly configured:\n\n%s\n\n" + "See docs/postgres.md for more information. You can override this check by" + "setting 'allow_unsafe_locale' to true in the database config.", + "\n".join(errors), + ) + + def lock_table(self, txn: Cursor, table: str) -> None: + txn.execute("LOCK TABLE %s in EXCLUSIVE MODE" % (table,)) diff --git a/synapse/storage/engines/postgres_rust.py b/synapse/storage/engines/postgres_rust.py index 82f4a91c58..7d5b100714 100644 --- a/synapse/storage/engines/postgres_rust.py +++ b/synapse/storage/engines/postgres_rust.py @@ -41,7 +41,8 @@ import logging from typing import TYPE_CHECKING, Any, Mapping from synapse.storage.engines._base import AUTO_INCREMENT_PRIMARY_KEYPLACEHOLDER -from synapse.storage.engines.postgres import PostgresEngine +from synapse.storage.engines.postgres_base import PostgresEngine +from synapse.storage.types import Connection, Cursor from synapse.synapse_rust.database import postgres if TYPE_CHECKING: @@ -53,15 +54,18 @@ logger = logging.getLogger(__name__) _RETRYABLE_PGCODES = ("40001", "40P01") -class RustPostgresEngine(PostgresEngine): +class RustPostgresEngine(PostgresEngine[Connection, Cursor]): """A :class:`PostgresEngine` that talks to the Rust backend's shim.""" def __init__(self, database_config: Mapping[str, Any]): - super().__init__(database_config) - # Route the DBAPI2 exception hierarchy (OperationalError, DatabaseError, - # IntegrityError, …) to the Rust backend's classes; the transaction - # driver catches `engine.module.`. - self.module = postgres + # The module is the Rust backend's DBAPI2 exception hierarchy + # (OperationalError, DatabaseError, IntegrityError, …); the transaction + # driver catches `engine.module.`. It is an intentionally *partial* + # `DBAPI2Module`: it exposes only the exception subset Synapse actually + # uses and has no module-level `connect` (connections come from the pool, + # via `rust_dbapi.connect`), so it doesn't structurally satisfy the + # protocol — hence the ignore. + super().__init__(postgres, database_config) # type: ignore[arg-type] def convert_param_style(self, sql: str) -> str: # The shim binds positional `$1, $2, ...` placeholders (like libpq), @@ -120,3 +124,20 @@ class RustPostgresEngine(PostgresEngine): "BIGINT PRIMARY KEY GENERATED ALWAYS AS IDENTITY", ) cursor.executescript(script) + + def check_database( + self, db_conn: Any, allow_outdated_version: bool = False + ) -> None: + # Startup database validation reads psycopg2 connection attributes + # (server_version, ...) that the shim doesn't expose; adapting it is + # part of wiring the Rust backend into startup (a follow-up). + raise NotImplementedError( + "check_database is not yet implemented for the Rust Postgres backend" + ) + + @property + def server_version(self) -> str: + # As above: depends on the psycopg2 startup path that isn't wired yet. + raise NotImplementedError( + "server_version is not yet implemented for the Rust Postgres backend" + ) diff --git a/tests/storage/test_room_search.py b/tests/storage/test_room_search.py index 2c0ef19e9e..a3ba7efb36 100644 --- a/tests/storage/test_room_search.py +++ b/tests/storage/test_room_search.py @@ -30,7 +30,7 @@ from synapse.rest.client import login, room from synapse.server import HomeServer from synapse.storage.databases.main import DataStore from synapse.storage.databases.main.search import Phrase, SearchToken, _tokenize_query -from synapse.storage.engines import PostgresEngine +from synapse.storage.engines import PostgresEngine, Psycopg2Engine from synapse.storage.engines.sqlite import Sqlite3Engine from synapse.util.clock import Clock @@ -273,7 +273,7 @@ class MessageSearchTest(HomeserverTestCase): # from ignoring the initial double quote to treating it as a phrase. main_store = homeserver.get_datastores().main found = False - if isinstance(main_store.database_engine, PostgresEngine): + if isinstance(main_store.database_engine, Psycopg2Engine): assert main_store.database_engine._version is not None found = main_store.database_engine._version < 140000 self.COMMON_CASES.append(('"fox quick', found))