From 39036ddf96894c8ead8f0d75bfc1b0641206f20f Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Tue, 8 Sep 2026 06:56:34 +0200 Subject: [PATCH 1/2] feat: generalize .prod.sql into a ..sql environment-tag convention testdb up/reset and migrate check/apply now accept --env to select which environment-tagged SQL files are in scope. testdb defaults to local_test; migrate's --env is optional and applies no filtering when omitted, matching today's behavior. Untagged files always apply. --- docs/database-layout.md | 14 ++++++++-- pgdevkit/cli.py | 29 ++++++++++++++----- pgdevkit/envtag.py | 54 ++++++++++++++++++++++++++++++++++++ pgdevkit/migrate.py | 13 +++++++-- pgdevkit/testdb/api.py | 20 ++++++++----- pgdevkit/testdb/mssql/api.py | 11 ++++---- pgdevkit/testdb/schema.py | 20 ++++++++----- 7 files changed, 130 insertions(+), 31 deletions(-) create mode 100644 pgdevkit/envtag.py diff --git a/docs/database-layout.md b/docs/database-layout.md index b27f709..95302d6 100644 --- a/docs/database-layout.md +++ b/docs/database-layout.md @@ -71,10 +71,18 @@ One object per file: `tables/user.sql`, `views/all_edits.sql`, |---|---| | `.sql` | The object's live definition (`CREATE TABLE`, `CREATE OR REPLACE VIEW`, ...) | | `.test_data.json` | Seed rows for a table — a JSON array of row objects, loaded after the table is created | -| `.init.sql` | One-time setup for an object (e.g. a backfill), run once, kept separate from the reusable definition | -| `.prod.sql` / `.prod` anywhere in the name | Production-only (real permission grants, real user accounts) — skipped by `pgdb testdb` | +| `.init.sql` | One-time setup for an object (e.g. a backfill), run once, kept separate from the reusable definition — `init` is reserved and is never treated as an environment tag | +| `..sql` | Only applied when targeting environment `` (any name you like — `prod`, `staging`, ...); a file with no such suffix is untagged and always applies, regardless of environment | | `all.sql` | Generated concatenation of the whole tree — not hand-edited, not committed | +`pgdb testdb up`/`pgdb testdb reset` apply the `--env` they're given (default +`local_test`) — so an untagged `grants.sql` always applies, but +`grants.prod.sql` is skipped unless you pass `--env prod`. `pgdb migrate +check`/`pgdb migrate apply` accept the same `--env`, but it's optional with no +default: omit it and every file is a candidate regardless of its tag (today's +behavior); pass it to restrict to files tagged for that environment plus +untagged ones. + --- ## Migrations @@ -156,5 +164,5 @@ leading sort number. - [ ] Object-type folder (`tables`, `views`, ...) matches the apply-order table above — that's what governs ordering, not the layer's leading number - [ ] One-off changes go in `migrations/`, dated, never edited after applying - [ ] The live `.sql` file is updated in the same change as any migration touching that object -- [ ] `.prod` files are production-only and skipped by `pgdb testdb` +- [ ] `..sql` files (e.g. `.prod.sql`) are skipped by `pgdb testdb` unless it's run with a matching `--env` - [ ] Every table (and non-obvious column) has a `COMMENT ON`, placed in the object's own `.sql` file diff --git a/pgdevkit/cli.py b/pgdevkit/cli.py index 9cdacdf..0fb2370 100644 --- a/pgdevkit/cli.py +++ b/pgdevkit/cli.py @@ -29,6 +29,17 @@ _EXCLUDE_AREA_OPTION = typer.Option( [], "--exclude-area", help="Skip files declaring this area (repeatable); untagged files are never excluded" ) +_TESTDB_ENV_OPTION = typer.Option( + "local_test", + "--env", + help="Environment to apply: skips any ..sql file (e.g. grants.prod.sql); untagged files always apply", +) +_MIGRATE_ENV_OPTION = typer.Option( + None, + "--env", + help="Restrict to files tagged for this environment (e.g. .prod.sql); untagged files always apply; " + "omit to apply every file regardless of its env tag", +) def _as_area_set(values: list[str]) -> frozenset[str] | None: @@ -188,17 +199,17 @@ def fetch_missing( @testdb_app.command("up") -def testdb_up() -> None: +def testdb_up(env: str = _TESTDB_ENV_OPTION) -> None: """Ensure the container is running, the workspace DB exists, and schema is applied.""" - testdb.ensure_testdb() + testdb.ensure_testdb(env=env) info = testdb.status() console.print(f"[green]Test DB ready:[/green] {info['database']} ({info['dsn']})") @testdb_app.command("reset") -def testdb_reset() -> None: +def testdb_reset(env: str = _TESTDB_ENV_OPTION) -> None: """Drop and recreate only this workspace's database, then reapply schema + seed data.""" - testdb.reset_testdb() + testdb.reset_testdb(env=env) info = testdb.status() console.print(f"[green]Test DB reset:[/green] {info['database']}") @@ -272,6 +283,7 @@ def migrate_check( ), area: list[str] = _AREA_OPTION, exclude_area: list[str] = _EXCLUDE_AREA_OPTION, + env: str | None = _MIGRATE_ENV_OPTION, ) -> None: """List which migration files under migrations_dir are applied vs. pending.""" if not migrations_dir.is_dir(): @@ -281,7 +293,7 @@ def migrate_check( conninfo = build_conninfo(url, entra_user) tracking_table = tracking_table or migrate.default_tracking_table(migrations_dir) local_files = migrate.list_migration_files( - migrations_dir, areas=_as_area_set(area), exclude_areas=_as_area_set(exclude_area) + migrations_dir, areas=_as_area_set(area), exclude_areas=_as_area_set(exclude_area), env=env ) try: applied = migrate.applied_migrations(conninfo, tracking_table) @@ -323,6 +335,7 @@ def migrate_apply( yes: bool = typer.Option(False, "--yes", "-y", help="Skip the confirm-target prompt"), area: list[str] = _AREA_OPTION, exclude_area: list[str] = _EXCLUDE_AREA_OPTION, + env: str | None = _MIGRATE_ENV_OPTION, ) -> None: """Apply pending migration files, in filename order, tracking each in tracking_table.""" if not migrations_dir.is_dir(): @@ -341,13 +354,15 @@ def migrate_apply( else: try: targets = migrate.pending_migrations( - migrations_dir, conninfo, tracking_table, areas=areas, exclude_areas=exclude_areas + migrations_dir, conninfo, tracking_table, areas=areas, exclude_areas=exclude_areas, env=env ) except migrate.TrackingTableMissing: err_console.print( f"[yellow]⚠[/yellow] {tracking_table} not found — treating every migration as pending" ) - targets = migrate.list_migration_files(migrations_dir, areas=areas, exclude_areas=exclude_areas) + targets = migrate.list_migration_files( + migrations_dir, areas=areas, exclude_areas=exclude_areas, env=env + ) if not targets: console.print("No pending migrations.") diff --git a/pgdevkit/envtag.py b/pgdevkit/envtag.py new file mode 100644 index 0000000..943eb6b --- /dev/null +++ b/pgdevkit/envtag.py @@ -0,0 +1,54 @@ +"""Optional `..sql` filename convention: a file whose dot-segment +immediately before `.sql` names a deployment environment (e.g. +`grants.prod.sql`, `seed.staging.sql`) is only in scope when the caller is +targeting that same environment. A plain `.sql` file (no such segment) +is untagged/common and is always in scope, regardless of which environment is +requested — mirroring the untagged-file rule for `-- area:` tags in +areas.py. + +This generalizes the older, hardcoded `.prod.sql` convention (still the usual +name for a production-only file — grants, real user accounts — that +`pgdb testdb` should never touch); any string can now be used as an +environment name. + +`.init.sql` (one-time setup, see docs/database-layout.md) is reserved and is +never interpreted as an environment tag. +""" + +from __future__ import annotations + +from pathlib import Path + +_RESERVED_SQL_SUFFIXES = {"init"} + + +def file_env(path: Path) -> str | None: + """The environment tag from `path`'s name, or None if it's untagged (or + the suffix is a reserved, non-env one like `.init.sql`). Only `.sql` + files can carry a tag.""" + if path.suffix != ".sql": + return None + stem = path.stem + base, dot, suffix = stem.rpartition(".") + if not dot or suffix in _RESERVED_SQL_SUFFIXES: + return None + return suffix + + +def env_allowed(path: Path, env: str | None) -> bool: + """Whether `path` is in scope for `env`. `env=None` means no environment + filtering was requested, so every file (tagged or not) is in scope.""" + if env is None: + return True + tag = file_env(path) + return tag is None or tag == env + + +def strip_env_suffix(path: Path) -> str: + """`path.stem` with a trailing `.` removed, so a tagged file + resolves to the same logical name as its untagged counterpart would + (e.g. `grants.prod.sql` -> "grants", same as `grants.sql`).""" + tag = file_env(path) + if tag is None: + return path.stem + return path.stem[: -(len(tag) + 1)] diff --git a/pgdevkit/migrate.py b/pgdevkit/migrate.py index ad0289e..6f3b543 100644 --- a/pgdevkit/migrate.py +++ b/pgdevkit/migrate.py @@ -19,6 +19,7 @@ from psycopg import sql as pg_sql from .areas import filter_by_area +from .envtag import env_allowed _IDENTIFIER = r"[A-Za-z_][A-Za-z0-9_]*" _DEFAULT_TRACKING_TABLE = "public.schema_migrations" @@ -281,9 +282,16 @@ def list_migration_files( *, areas: frozenset[str] | None = None, exclude_areas: frozenset[str] | None = None, + env: str | None = None, ) -> list[Path]: + """Migration files under migrations_dir, restricted by area (see + `.areas`) and by environment tag (see `.envtag`) — e.g. a + `2026-07-10_backfill.prod.sql` is only included when `env="prod"`. + `env=None` (the default) applies no environment filtering at all, so + every file is a candidate regardless of its tag.""" files = sorted(migrations_dir.glob("*.sql")) - return filter_by_area(files, only=areas, exclude=exclude_areas) + files = filter_by_area(files, only=areas, exclude=exclude_areas) + return [f for f in files if env_allowed(f, env)] def applied_migrations(conninfo: str, tracking_table: str) -> dict[str, tuple[datetime, str]]: @@ -306,9 +314,10 @@ def pending_migrations( *, areas: frozenset[str] | None = None, exclude_areas: frozenset[str] | None = None, + env: str | None = None, ) -> list[Path]: applied = applied_migrations(conninfo, tracking_table) - files = list_migration_files(migrations_dir, areas=areas, exclude_areas=exclude_areas) + files = list_migration_files(migrations_dir, areas=areas, exclude_areas=exclude_areas, env=env) return [p for p in files if p.name not in applied] diff --git a/pgdevkit/testdb/api.py b/pgdevkit/testdb/api.py index 2e0e85a..12b5073 100644 --- a/pgdevkit/testdb/api.py +++ b/pgdevkit/testdb/api.py @@ -68,24 +68,30 @@ async def _drop_database(db_name: str) -> None: await con.execute(SQL("DROP DATABASE IF EXISTS {}").format(Identifier(db_name))) -async def _apply(config: ProjectConfig, db_name: str, force_reset: bool) -> None: +async def _apply(config: ProjectConfig, db_name: str, force_reset: bool, env: str) -> None: async with await psycopg.AsyncConnection.connect(_db_dsn(db_name), autocommit=True) as con: await apply_schema( con, config.root / config.database_dir, extensions=config.extensions, force_reset=force_reset, + env=env, ) -def ensure_testdb(project_root: Path | None = None, force_reset: bool = False) -> dict[str, str]: +def ensure_testdb( + project_root: Path | None = None, force_reset: bool = False, env: str = "local_test" +) -> dict[str, str]: """Ensure the shared container is running, this workspace's database exists, and its schema is applied. Returns the {PREFIX}POSTGRES_* env vars for this workspace (or the mssql equivalent's env vars, per - `config.engine`).""" + `config.engine`). + + `env` selects which environment-tagged files apply (see pgdevkit.envtag, + e.g. a `grants.prod.sql` is skipped unless env="prod").""" config, db_name = _resolve(project_root) if config.engine == "mssql": - return _mssql_api().ensure_testdb(config, db_name, force_reset) + return _mssql_api().ensure_testdb(config, db_name, force_reset, env) ensure_container() @@ -93,16 +99,16 @@ async def _run() -> None: if force_reset: await _drop_database(db_name) await _ensure_database(db_name) - await _apply(config, db_name, force_reset) + await _apply(config, db_name, force_reset, env) asyncio.run(_run()) return _env_for(config, db_name) -def reset_testdb(project_root: Path | None = None) -> dict[str, str]: +def reset_testdb(project_root: Path | None = None, env: str = "local_test") -> dict[str, str]: """Drop and recreate only this workspace's database, then reapply schema and seed data.""" - return ensure_testdb(project_root, force_reset=True) + return ensure_testdb(project_root, force_reset=True, env=env) def clean_testdb(project_root: Path | None = None, all: bool = False) -> None: diff --git a/pgdevkit/testdb/mssql/api.py b/pgdevkit/testdb/mssql/api.py index 40af093..55dbb89 100644 --- a/pgdevkit/testdb/mssql/api.py +++ b/pgdevkit/testdb/mssql/api.py @@ -9,6 +9,7 @@ from ...db.mssql_sql import ident, json_encode_value from ...dialect import MSSQL +from ...envtag import strip_env_suffix from .. import query from ..config import ProjectConfig from ..schema import _iter_sql_files, _strip_layer_prefix @@ -101,7 +102,7 @@ def _run() -> None: await asyncio.to_thread(_run) -async def _apply(config: ProjectConfig, db_name: str, force_reset: bool) -> None: +async def _apply(config: ProjectConfig, db_name: str, force_reset: bool, env: str) -> None: def _connect() -> Any: return mssql_python.connect(_db_dsn(db_name), autocommit=True) @@ -110,7 +111,7 @@ def _connect() -> Any: database_dir = config.root / config.database_dir if not database_dir.is_dir(): return - for file, sql in _iter_sql_files(database_dir, MSSQL): + for file, sql in _iter_sql_files(database_dir, MSSQL, env): for batch in query.split_tsql_batches(sql): def _exec(batch: str = batch) -> None: @@ -120,20 +121,20 @@ def _exec(batch: str = batch) -> None: json_file = file.with_suffix(".test_data.json") if json_file.exists(): schema_name = _strip_layer_prefix(file.parent.parent.name) - table_stem = _strip_layer_prefix(file.stem) + table_stem = _strip_layer_prefix(strip_env_suffix(file)) await _insert_test_data(json_file, f"{schema_name}.{table_stem}", force_reset, conn) finally: await asyncio.to_thread(conn.close) -def ensure_testdb(config: ProjectConfig, db_name: str, force_reset: bool) -> dict[str, str]: +def ensure_testdb(config: ProjectConfig, db_name: str, force_reset: bool, env: str = "local_test") -> dict[str, str]: ensure_mssql_container() async def _run() -> None: if force_reset: await _drop_database(db_name) await _ensure_database(db_name) - await _apply(config, db_name, force_reset) + await _apply(config, db_name, force_reset, env) asyncio.run(_run()) return _env_for(config, db_name) diff --git a/pgdevkit/testdb/schema.py b/pgdevkit/testdb/schema.py index 581e08b..ea0cc7b 100644 --- a/pgdevkit/testdb/schema.py +++ b/pgdevkit/testdb/schema.py @@ -15,6 +15,7 @@ from ..db.complex_types import ComplexHelper from ..dialect import Dialect, POSTGRES +from ..envtag import env_allowed, strip_env_suffix from ..parser import IGNORED_DIR_NAMES logger = logging.getLogger(__name__) @@ -38,7 +39,7 @@ def _get_type_order(path: Path) -> int: - filename = re.sub(r"^\d+(\.\d+)?", "", path.name).removeprefix("_").removesuffix(".sql") + filename = re.sub(r"^\d+(\.\d+)?", "", strip_env_suffix(path)).removeprefix("_") if filename in _TYPE_ORDER: return _TYPE_ORDER[filename] if path.parent.name in _TYPE_ORDER: @@ -119,7 +120,7 @@ def _get_sql_deps(sql: str, dialect: Dialect = POSTGRES) -> set[str]: return deps -def _iter_sql_files(database_dir: Path, dialect: Dialect = POSTGRES): +def _iter_sql_files(database_dir: Path, dialect: Dialect = POSTGRES, env: str = "local_test"): """Yield (Path, sql_content) pairs in dependency-safe execution order.""" files: list[Path] = [] for root, dirs, dbfiles in os.walk(database_dir): @@ -129,8 +130,9 @@ def _iter_sql_files(database_dir: Path, dialect: Dialect = POSTGRES): for file in dbfiles: if file in ("all.sql", "100_permissions.sql"): continue - if file.endswith(".sql") and ".prod" not in file: - files.append(Path(root) / file) + path = Path(root) / file + if file.endswith(".sql") and env_allowed(path, env): + files.append(path) delivered: set[str] = set() delayed: list[tuple[str | None, Path, str]] = [] @@ -141,7 +143,7 @@ def _iter_sql_files(database_dir: Path, dialect: Dialect = POSTGRES): deps = _get_sql_deps(content, dialect) if file.parent.name in _SCHEMA_QUALIFIED_TYPES: schema = _strip_layer_prefix(file.parent.parent.name) - full_name = f"{schema}.{_strip_layer_prefix(file.stem)}" + full_name = f"{schema}.{_strip_layer_prefix(strip_env_suffix(file))}" deps.discard(full_name) # the file's own CREATE target is not a real dependency all_declared.add(full_name) if not deps or all(d in delivered for d in deps): @@ -242,10 +244,14 @@ async def apply_schema( force_reset: bool = False, *, dialect: Dialect = POSTGRES, + env: str = "local_test", ) -> None: """Apply every .sql file under database_dir (in dependency-safe order) and seed any matching .test_data.json files. Safe to call repeatedly. + A file tagged for another environment (e.g. `grants.prod.sql` when + `env="local_test"`) is skipped — see pgdevkit.envtag. + `migrations/` subdirectories are never applied here — they're for one-time manual application against real (already-provisioned) databases, not for building a fresh schema. The base object files under @@ -261,11 +267,11 @@ async def _apply(file: Path, sql: str) -> None: json_file = file.with_suffix(".test_data.json") if json_file.exists(): schema_name = _strip_layer_prefix(file.parent.parent.name) - table_stem = _strip_layer_prefix(file.stem) + table_stem = _strip_layer_prefix(strip_env_suffix(file)) await _insert_test_data(json_file, f"{schema_name}.{table_stem}", force_reset, con, complex_helper) failures: list[tuple[Path, str]] = [] - for file, sql in _iter_sql_files(database_dir, dialect): + for file, sql in _iter_sql_files(database_dir, dialect, env): try: await _apply(file, sql) except Exception as e: # noqa: BLE001 From 909acf9e321cb0d39068c23e426c74f858c82e86 Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Tue, 8 Sep 2026 07:00:56 +0200 Subject: [PATCH 2/2] test: cover env-tag filtering in envtag/migrate/testdb schema apply, docs Adds unit tests for pgdevkit.envtag, migrate list_migration_files/pending_migrations env filtering, and integration tests for apply_schema's env-tagged file handling. Documents the ..sql convention in README, SKILL.md, and database-layout.md. --- README.md | 29 +++++++++- skills/pgdevkit/SKILL.md | 4 +- tests/test_envtag.py | 56 +++++++++++++++++++ tests/test_migrate_env.py | 44 +++++++++++++++ .../database/app/tables/prod_only.prod.sql | 3 + .../database/app/tables/widget.init.sql | 1 + tests/testdb/test_schema.py | 37 ++++++++++++ 7 files changed, 171 insertions(+), 3 deletions(-) create mode 100644 tests/test_envtag.py create mode 100644 tests/test_migrate_env.py create mode 100644 tests/testdb/fixtures/database/app/tables/prod_only.prod.sql create mode 100644 tests/testdb/fixtures/database/app/tables/widget.init.sql diff --git a/README.md b/README.md index c72e2fe..af2259a 100644 --- a/README.md +++ b/README.md @@ -114,6 +114,32 @@ would then reconstruct a duplicate file for something that already exists. keyword arguments (`pgdevkit.fetch_missing.find_missing_objects` doesn't, for the reason above). +## Environment-tagged files (`..sql`) + +A file whose name ends `..sql` (e.g. `grants.prod.sql`, +`seed.staging.sql`) is only in scope when targeting that environment; a +plain `.sql` file is untagged and always in scope, regardless of +environment. `.init.sql` (see `docs/database-layout.md`) is reserved and is +never treated as an environment tag. + +- `pgdb testdb up`/`pgdb testdb reset` accept `--env` (default + `local_test`) — so an untagged `grants.sql` always applies, but + `grants.prod.sql` is skipped unless run with `--env prod`. +- `pgdb migrate check`/`pgdb migrate apply` accept `--env` too, but it's + optional with **no** default: omit it and every file is a candidate + regardless of its tag (unchanged, today's behavior); pass it to restrict + to files tagged for that environment plus untagged ones. + +```bash +pgdb testdb up --env prod # apply prod-tagged files too, against the local test container +pgdb migrate apply path/to/database/_migration_scripts --url ... --env prod +``` + +`pgdevkit.envtag` exposes the same logic for scripting: `file_env` reads a +file's tag, `env_allowed` applies the filtering semantics above, and +`strip_env_suffix` returns a tagged file's logical name (e.g. +`grants.prod.sql` -> `"grants"`). + ## `pgdb testdb` Manages a single shared, Podman-backed Postgres container for local tests @@ -142,7 +168,8 @@ def ensure_test_postgres(): os.environ[k] = v ``` -CLI: `pgdb testdb up|reset|run-sql|status|shell|clean`. +CLI: `pgdb testdb up|reset|run-sql|status|shell|clean`. `up`/`reset` accept +`--env` (default `local_test`) — see "Environment-tagged files" above. Container connection defaults (`localhost:54322`, `postgres`/`testpwd`) can be overridden with `PGDEVKIT_TESTDB_HOST`, `PGDEVKIT_TESTDB_PORT`, diff --git a/skills/pgdevkit/SKILL.md b/skills/pgdevkit/SKILL.md index 8c0e03a..f0866f1 100644 --- a/skills/pgdevkit/SKILL.md +++ b/skills/pgdevkit/SKILL.md @@ -253,7 +253,7 @@ sqlfmt --check db/queries/ # CI check ## The `database/` folder & backfilling untracked objects -See [docs/database-layout.md](../../docs/database-layout.md) for the full convention: layer directories, object-type subfolders and their apply order, file-naming rules (`.test_data.json`, `.init.sql`, `.prod`), and how migrations are organised. +See [docs/database-layout.md](../../docs/database-layout.md) for the full convention: layer directories, object-type subfolders and their apply order, file-naming rules (`.test_data.json`, `.init.sql`, `..sql`), and how migrations are organised. If a table, view, or function was created directly on the database and never got a `.sql` file: @@ -285,6 +285,6 @@ Reports drift between the `database/` `.sql` files and the actual schema — tab - [ ] All parameters use `%(name)s` style with a dict argument - [ ] Results mapped to a Pydantic model; table-mapped models extend `PostgresTableModel` - [ ] No `LATERAL JOIN` — use a CTE that groups/aggregates first, then joins it -- [ ] `.prod` files are production-only and skipped by `pgdb testdb` +- [ ] `..sql` files (e.g. `.prod.sql`) are skipped by `pgdb testdb` unless it's run with a matching `--env` - [ ] Every table (and non-obvious column) has a `COMMENT ON`, placed in the object's own `.sql` file - [ ] Untracked DB objects backfilled via `pgdb fetch-missing`, not left undocumented diff --git a/tests/test_envtag.py b/tests/test_envtag.py new file mode 100644 index 0000000..aa0dfda --- /dev/null +++ b/tests/test_envtag.py @@ -0,0 +1,56 @@ +from __future__ import annotations + +from pathlib import Path + +from pgdevkit.envtag import env_allowed, file_env, strip_env_suffix + + +class TestFileEnv: + def test_plain_sql_file_is_untagged(self): + assert file_env(Path("grants.sql")) is None + + def test_init_sql_is_never_a_tag(self): + assert file_env(Path("user.init.sql")) is None + + def test_prod_suffix_is_the_prod_tag(self): + assert file_env(Path("grants.prod.sql")) == "prod" + + def test_arbitrary_env_name_is_a_tag(self): + assert file_env(Path("seed.staging.sql")) == "staging" + assert file_env(Path("seed.local_test.sql")) == "local_test" + + def test_non_sql_file_is_never_tagged(self): + assert file_env(Path("widget.test_data.json")) is None + + +class TestEnvAllowed: + def test_no_env_requested_allows_everything(self): + assert env_allowed(Path("grants.prod.sql"), None) is True + assert env_allowed(Path("grants.sql"), None) is True + + def test_untagged_file_always_allowed(self): + assert env_allowed(Path("grants.sql"), "prod") is True + assert env_allowed(Path("grants.sql"), "staging") is True + + def test_matching_tag_allowed(self): + assert env_allowed(Path("grants.prod.sql"), "prod") is True + + def test_mismatched_tag_disallowed(self): + assert env_allowed(Path("grants.prod.sql"), "staging") is False + assert env_allowed(Path("grants.prod.sql"), "local_test") is False + + def test_init_file_always_allowed(self): + assert env_allowed(Path("user.init.sql"), "prod") is True + assert env_allowed(Path("user.init.sql"), "staging") is True + + +class TestStripEnvSuffix: + def test_untagged_file_unchanged(self): + assert strip_env_suffix(Path("grants.sql")) == "grants" + + def test_tagged_file_matches_untagged_logical_name(self): + assert strip_env_suffix(Path("grants.prod.sql")) == "grants" + assert strip_env_suffix(Path("seed.local_test.sql")) == "seed" + + def test_init_file_keeps_init_in_the_stem(self): + assert strip_env_suffix(Path("user.init.sql")) == "user.init" diff --git a/tests/test_migrate_env.py b/tests/test_migrate_env.py new file mode 100644 index 0000000..c76939a --- /dev/null +++ b/tests/test_migrate_env.py @@ -0,0 +1,44 @@ +from __future__ import annotations + +from pathlib import Path + +from pgdevkit.migrate import list_migration_files + + +def _write(dir: Path, name: str, content: str = "select 1;\n") -> Path: + p = dir / name + p.write_text(content, encoding="utf-8") + return p + + +class TestListMigrationFilesEnvFiltering: + def test_no_env_applies_every_file_regardless_of_tag(self, tmp_path: Path): + prod = _write(tmp_path, "001_prod.prod.sql") + common = _write(tmp_path, "002_common.sql") + + assert list_migration_files(tmp_path) == sorted([prod, common]) + + def test_env_keeps_matching_tag_and_untagged(self, tmp_path: Path): + prod = _write(tmp_path, "001_backfill.prod.sql") + staging = _write(tmp_path, "002_backfill.staging.sql") + common = _write(tmp_path, "003_common.sql") + + result = list_migration_files(tmp_path, env="prod") + assert set(result) == {prod, common} + assert staging not in result + + def test_init_file_is_never_filtered_by_env(self, tmp_path: Path): + init = _write(tmp_path, "001_setup.init.sql") + + assert list_migration_files(tmp_path, env="prod") == [init] + assert list_migration_files(tmp_path, env="staging") == [init] + + def test_env_and_area_filters_compose(self, tmp_path: Path): + billing_prod = _write(tmp_path, "001_billing.prod.sql", "-- area: billing\nselect 1;\n") + billing_staging = _write(tmp_path, "002_billing.staging.sql", "-- area: billing\nselect 1;\n") + reporting_prod = _write(tmp_path, "003_reporting.prod.sql", "-- area: reporting\nselect 1;\n") + + result = list_migration_files(tmp_path, areas=frozenset({"billing"}), env="prod") + assert result == [billing_prod] + assert billing_staging not in result + assert reporting_prod not in result diff --git a/tests/testdb/fixtures/database/app/tables/prod_only.prod.sql b/tests/testdb/fixtures/database/app/tables/prod_only.prod.sql new file mode 100644 index 0000000..6545246 --- /dev/null +++ b/tests/testdb/fixtures/database/app/tables/prod_only.prod.sql @@ -0,0 +1,3 @@ +CREATE TABLE IF NOT EXISTS app.prod_only ( + id serial PRIMARY KEY +); diff --git a/tests/testdb/fixtures/database/app/tables/widget.init.sql b/tests/testdb/fixtures/database/app/tables/widget.init.sql new file mode 100644 index 0000000..295bd64 --- /dev/null +++ b/tests/testdb/fixtures/database/app/tables/widget.init.sql @@ -0,0 +1 @@ +ALTER TABLE app.widget ADD COLUMN IF NOT EXISTS init_flag boolean NOT NULL DEFAULT true; diff --git a/tests/testdb/test_schema.py b/tests/testdb/test_schema.py index bf7c8fe..eaa1c7e 100644 --- a/tests/testdb/test_schema.py +++ b/tests/testdb/test_schema.py @@ -114,6 +114,43 @@ async def test_apply_schema_seeds_composite_enum_and_jsonb_columns(schema_test_d assert tags == ["small", "shiny"] +@requires_podman +async def test_apply_schema_skips_env_tagged_file_for_a_different_env(schema_test_db): + # prod_only.prod.sql is tagged for "prod" -- with the default env + # ("local_test"), it must not be applied. + async with await psycopg.AsyncConnection.connect(_db_dsn(), autocommit=True) as con: + await apply_schema(con, FIXTURES) + async with con.cursor() as cur: + await cur.execute("SELECT to_regclass('app.prod_only')") + (regclass,) = await cur.fetchone() + assert regclass is None + + +@requires_podman +async def test_apply_schema_applies_env_tagged_file_for_the_matching_env(schema_test_db): + async with await psycopg.AsyncConnection.connect(_db_dsn(), autocommit=True) as con: + await apply_schema(con, FIXTURES, env="prod") + async with con.cursor() as cur: + await cur.execute("SELECT to_regclass('app.prod_only')") + (regclass,) = await cur.fetchone() + assert regclass is not None + + +@requires_podman +async def test_apply_schema_init_file_applies_regardless_of_env(schema_test_db): + # widget.init.sql adds init_flag to app.widget -- "init" is a reserved + # suffix, never an environment tag, so this must apply under any env. + async with await psycopg.AsyncConnection.connect(_db_dsn(), autocommit=True) as con: + await apply_schema(con, FIXTURES, env="staging") + async with con.cursor() as cur: + await cur.execute( + "SELECT column_name FROM information_schema.columns " + "WHERE table_schema = 'app' AND table_name = 'widget' AND column_name = 'init_flag'" + ) + row = await cur.fetchone() + assert row is not None + + @requires_podman async def test_apply_schema_never_applies_migrations_dir(schema_test_db): # migrations/ is for one-time manual application against real databases,