From dda93db0acd3afce6aba4b32298afc4125a196be Mon Sep 17 00:00:00 2001 From: tegwick Date: Tue, 1 Sep 2026 02:05:56 +0200 Subject: [PATCH] fix: commit migration session setup Assistant: codex Assistant-Model: gpt-5.6-sol Assistant-Session: 01a053ff-1d6f-7fe2-ac1c-a6eb40a42a0c --- hub_core/migrations/env.py | 11 +++++----- hub_core/migrations/roles.py | 31 ++++++++++++++++++++++++++++ tests/test_migration_environment.py | 32 ++++++++++++++++++++++++++++- 3 files changed, 68 insertions(+), 6 deletions(-) diff --git a/hub_core/migrations/env.py b/hub_core/migrations/env.py index 8d4a970..0031192 100644 --- a/hub_core/migrations/env.py +++ b/hub_core/migrations/env.py @@ -4,8 +4,8 @@ from logging.config import fileConfig from alembic import context from sqlalchemy import engine_from_config, pool +from hub_core.migrations.roles import configure_migration_session from hub_core.models import Base -from hub_core.migrations.roles import migration_role_statement, migration_schema_statement config = context.config @@ -39,10 +39,11 @@ def run_migrations_online() -> None: poolclass=pool.NullPool, ) with connectable.connect() as connection: - if statement := migration_role_statement(os.environ.get("HUB_CORE_MIGRATION_ROLE")): - connection.exec_driver_sql(statement) - if statement := migration_schema_statement(os.environ.get("HUB_CORE_MIGRATION_SCHEMA")): - connection.exec_driver_sql(statement) + configure_migration_session( + connection, + role=os.environ.get("HUB_CORE_MIGRATION_ROLE"), + schema=os.environ.get("HUB_CORE_MIGRATION_SCHEMA"), + ) context.configure(connection=connection, target_metadata=target_metadata) with context.begin_transaction(): context.run_migrations() diff --git a/hub_core/migrations/roles.py b/hub_core/migrations/roles.py index 4eb38c9..427658a 100644 --- a/hub_core/migrations/roles.py +++ b/hub_core/migrations/roles.py @@ -1,6 +1,13 @@ from __future__ import annotations import re +from typing import Protocol + + +class MigrationConnection(Protocol): + def exec_driver_sql(self, statement: str) -> object: ... + + def commit(self) -> None: ... def migration_role_statement(role: str | None) -> str | None: @@ -19,3 +26,27 @@ def migration_schema_statement(schema: str | None) -> str | None: if not re.fullmatch(r"[a-z_][a-z0-9_]{0,62}", schema): raise ValueError("HUB_CORE_MIGRATION_SCHEMA is not a safe PostgreSQL schema name") return f'SET search_path TO "{schema}", public' + + +def configure_migration_session( + connection: MigrationConnection, + *, + role: str | None, + schema: str | None, +) -> None: + """Apply and commit session settings before Alembic opens its transaction.""" + statements = tuple( + statement + for statement in ( + migration_role_statement(role), + migration_schema_statement(schema), + ) + if statement is not None + ) + for statement in statements: + connection.exec_driver_sql(statement) + if statements: + # SQLAlchemy 2 autobegins on exec_driver_sql(). Without this commit, + # Alembic joins the setup transaction and its DDL is rolled back when + # the connection closes. + connection.commit() diff --git a/tests/test_migration_environment.py b/tests/test_migration_environment.py index 93d055c..b1a0eee 100644 --- a/tests/test_migration_environment.py +++ b/tests/test_migration_environment.py @@ -1,8 +1,14 @@ from __future__ import annotations +from unittest.mock import Mock, call + import pytest -from hub_core.migrations.roles import migration_role_statement, migration_schema_statement +from hub_core.migrations.roles import ( + configure_migration_session, + migration_role_statement, + migration_schema_statement, +) def test_migration_role_is_quoted_and_validated() -> None: @@ -19,3 +25,27 @@ def test_migration_schema_is_explicitly_quoted_and_validated() -> None: assert migration_schema_statement(None) is None with pytest.raises(ValueError, match="safe PostgreSQL schema"): migration_schema_statement('hub_runtime"; DROP SCHEMA public; --') + + +def test_migration_session_commits_setup_before_alembic_transaction() -> None: + connection = Mock() + + configure_migration_session( + connection, + role="hub_runtime_owner", + schema="hub_runtime", + ) + + assert connection.method_calls == [ + call.exec_driver_sql('SET ROLE "hub_runtime_owner"'), + call.exec_driver_sql('SET search_path TO "hub_runtime", public'), + call.commit(), + ] + + +def test_migration_session_does_not_open_empty_setup_transaction() -> None: + connection = Mock() + + configure_migration_session(connection, role=None, schema=None) + + connection.assert_not_called()