diff --git a/hub_core/migrations/env.py b/hub_core/migrations/env.py index d322e7b..8d4a970 100644 --- a/hub_core/migrations/env.py +++ b/hub_core/migrations/env.py @@ -5,7 +5,7 @@ from alembic import context from sqlalchemy import engine_from_config, pool from hub_core.models import Base -from hub_core.migrations.roles import migration_role_statement +from hub_core.migrations.roles import migration_role_statement, migration_schema_statement config = context.config @@ -41,6 +41,8 @@ def run_migrations_online() -> None: 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) 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 7f93ae3..4eb38c9 100644 --- a/hub_core/migrations/roles.py +++ b/hub_core/migrations/roles.py @@ -10,3 +10,12 @@ def migration_role_statement(role: str | None) -> str | None: if not re.fullmatch(r"[a-z_][a-z0-9_]{0,62}", role): raise ValueError("HUB_CORE_MIGRATION_ROLE is not a safe PostgreSQL role name") return f'SET ROLE "{role}"' + + +def migration_schema_statement(schema: str | None) -> str | None: + """Return an explicit, safely quoted migration search path.""" + if not schema: + return 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' diff --git a/tests/test_migration_environment.py b/tests/test_migration_environment.py index 8eec3b4..93d055c 100644 --- a/tests/test_migration_environment.py +++ b/tests/test_migration_environment.py @@ -2,7 +2,7 @@ from __future__ import annotations import pytest -from hub_core.migrations.roles import migration_role_statement +from hub_core.migrations.roles import migration_role_statement, migration_schema_statement def test_migration_role_is_quoted_and_validated() -> None: @@ -10,3 +10,12 @@ def test_migration_role_is_quoted_and_validated() -> None: assert migration_role_statement(None) is None with pytest.raises(ValueError, match="safe PostgreSQL role"): migration_role_statement('owner"; DROP SCHEMA public; --') + + +def test_migration_schema_is_explicitly_quoted_and_validated() -> None: + assert migration_schema_statement("hub_runtime") == ( + 'SET search_path TO "hub_runtime", public' + ) + assert migration_schema_statement(None) is None + with pytest.raises(ValueError, match="safe PostgreSQL schema"): + migration_schema_statement('hub_runtime"; DROP SCHEMA public; --')