From a5843c6a6aef0754fdcf1efa0affb15b21a8408c Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Mon, 7 Sep 2026 20:31:33 +0200 Subject: [PATCH 1/5] Add pgdevkit.schemas: filter files by referenced DB schema Composable with the existing -- area: tag filtering, but derived automatically by parsing each file's SQL instead of requiring a tag. Wired into pgdb compare, migrate check/apply, and testdb up/reset (the latter two also gain --area/--exclude-area, which they lacked entirely until now). Co-Authored-By: Claude Sonnet 5 Claude-Session: https://claude.ai/code/session_01StWgZw5TDd7KNGY3VzNGMg --- README.md | 83 +++++++++++----- pgdevkit/cli.py | 82 +++++++++++++--- pgdevkit/migrate.py | 12 ++- pgdevkit/parser.py | 11 ++- pgdevkit/schemas.py | 112 ++++++++++++++++++++++ pgdevkit/testdb/api.py | 56 +++++++++-- pgdevkit/testdb/mssql/api.py | 32 ++++++- pgdevkit/testdb/schema.py | 45 ++++++++- tests/test_migrate_schemas.py | 52 +++++++++++ tests/test_parser_schemas.py | 44 +++++++++ tests/test_schemas.py | 130 ++++++++++++++++++++++++++ tests/testdb/test_schema_filtering.py | 62 ++++++++++++ 12 files changed, 666 insertions(+), 55 deletions(-) create mode 100644 pgdevkit/schemas.py create mode 100644 tests/test_migrate_schemas.py create mode 100644 tests/test_parser_schemas.py create mode 100644 tests/test_schemas.py create mode 100644 tests/testdb/test_schema_filtering.py diff --git a/README.md b/README.md index c72e2fe..abf978f 100644 --- a/README.md +++ b/README.md @@ -58,7 +58,13 @@ which parses/introspects/diffs like any other column type; see handled on the CRUD side (write-side serialization only, no auto-parsing on read — `mssql-python` doesn't distinguish `json` columns from `nvarchar`). -### Area tagging and filtering +### Area and schema filtering + +Two independent, composable ways to narrow which files a command touches: +**area** is an explicit opt-in tag; **schema** is derived automatically from +each file's own SQL. + +#### Area tagging Any migration file or `database/` code file can declare one or more areas by starting with a `-- area:` comment: @@ -76,43 +82,69 @@ first real statement); a `-- area:` comment later in the file doesn't count. A file with no directive is untagged, and untagged files are treated as shared/common. -`pgdb compare`, `pgdb migrate check`, and `pgdb migrate apply` all accept: +#### Schema filtering + +No tag needed — schema membership is parsed straight out of the SQL itself: +every schema-qualified (or default-schema, when unqualified) table/view/ +function/index reference across every statement in the file, DDL or DML +alike, plus any `CREATE SCHEMA name`. A file whose schema(s) can't be +determined (unparseable content, or no table/schema reference in it at all) +is treated the same as an untagged file — always kept. + +#### Options + +`pgdb compare`, `pgdb migrate check`, `pgdb migrate apply`, `pgdb testdb up`, +and `pgdb testdb reset` all accept: - `--area NAME` (repeatable) — restrict to files declaring one of the given areas, **plus every untagged file** (untagged files always stay in scope). - `--exclude-area NAME` (repeatable) — drop files declaring one of the given areas; untagged files are never dropped by this. +- `--schema NAME` (repeatable) — restrict to files referencing one of the + given schemas, **plus every file with no detectable schema reference**. +- `--exclude-schema NAME` (repeatable) — drop files referencing one of the + given schemas; files with no detectable reference are never dropped. -Both can be combined; a file matching both an included and an excluded area -is excluded. Passing neither option applies no filtering (the default, -unchanged behavior). +All four can be combined — a file must pass every filter it's subject to (an +area match doesn't excuse a schema mismatch, and vice versa), and a file +matching both an included and an excluded value on the same axis is +excluded. Passing none of them applies no filtering (the default, unchanged +behavior). ```bash pgdb migrate apply path/to/database/_migration_scripts --url ... --area billing pgdb compare path/to/database/ --url ... --exclude-area reporting +pgdb migrate check path/to/database/_migration_scripts --url ... --schema billing --exclude-schema reporting +pgdb testdb up --schema billing ``` `compare`'s default report (no `--report-extra-db`) only checks that the filtered scripts exist correctly in the DB, so it composes safely with area -filtering. Passing `--report-extra-db` together with an area filter also -reports every DB object outside the filtered area(s) as "missing in -scripts" — since the live database has no concept of areas, only the -scripts side is filtered — so treat that combination's "missing in scripts" -results with that in mind (the CLI prints a warning when you combine them). - -`pgdb fetch-missing` deliberately has **no** `--area`/`--exclude-area`: it -diffs the full database against scripts to find genuinely untracked -objects, so narrowing the scripts side by area would make every object -tracked only under a different area look "missing" too — and `--write` -would then reconstruct a duplicate file for something that already exists. - -`pgdevkit.areas` exposes the same logic for scripting: +and schema filtering. Passing `--report-extra-db` together with either kind +of filter also reports every DB object outside the filtered area(s)/ +schema(s) as "missing in scripts" — since the live database has no concept +of areas, and isn't itself filtered by `--schema` either — only the scripts +side is filtered — so treat that combination's "missing in scripts" results +with that in mind (the CLI prints a warning when you combine them). + +`pgdb fetch-missing` deliberately has **no** `--area`/`--exclude-area` (or +`--schema`/`--exclude-schema`): it diffs the full database against scripts +to find genuinely untracked objects, so narrowing the scripts side would +make every object tracked under a different area/schema look "missing" too +— and `--write` would then reconstruct a duplicate file for something that +already exists. + +`pgdevkit.areas` exposes the tag-filtering logic for scripting: `parse_areas`/`file_areas` read a file's declared areas, and `area_allowed`/`filter_by_area` apply the `only`/`exclude` semantics above. -`pgdevkit.migrate.list_migration_files`/`pending_migrations` and -`pgdevkit.parser.parse_directory` take the same `areas`/`exclude_areas` -keyword arguments (`pgdevkit.fetch_missing.find_missing_objects` doesn't, -for the reason above). +`pgdevkit.schemas` exposes the equivalent for schema filtering: +`sql_schemas`/`file_schemas` detect a file's referenced schemas, and +`schema_allowed`/`filter_by_schema` apply the same `only`/`exclude` +semantics. `pgdevkit.migrate.list_migration_files`/`pending_migrations` and +`pgdevkit.parser.parse_directory` take both pairs of keyword arguments +(`areas`/`exclude_areas` and `schemas`/`exclude_schemas`); +`pgdevkit.fetch_missing.find_missing_objects` takes neither, for the reason +above. ## `pgdb testdb` @@ -144,6 +176,13 @@ def ensure_test_postgres(): CLI: `pgdb testdb up|reset|run-sql|status|shell|clean`. +`up`/`reset` accept `--area`/`--exclude-area` and `--schema`/`--exclude-schema` +(see "Area and schema filtering" above) to scope which `database/` files get +applied — e.g. `pgdb testdb up --schema billing` for a test DB with only the +`billing` schema's tables/views/functions, without waiting on the rest of the +project's schema to apply. `ensure_testdb`/`reset_testdb` take the same +keyword arguments when called from Python (e.g. from a pytest fixture). + Container connection defaults (`localhost:54322`, `postgres`/`testpwd`) can be overridden with `PGDEVKIT_TESTDB_HOST`, `PGDEVKIT_TESTDB_PORT`, `PGDEVKIT_TESTDB_USER`, `PGDEVKIT_TESTDB_PASSWORD`. Before touching the diff --git a/pgdevkit/cli.py b/pgdevkit/cli.py index 9cdacdf..ea300ac 100644 --- a/pgdevkit/cli.py +++ b/pgdevkit/cli.py @@ -29,9 +29,21 @@ _EXCLUDE_AREA_OPTION = typer.Option( [], "--exclude-area", help="Skip files declaring this area (repeatable); untagged files are never excluded" ) +_SCHEMA_OPTION = typer.Option( + [], + "--schema", + help="Restrict to files referencing this DB schema (repeatable); " + "files with no detectable schema reference always stay in scope", +) +_EXCLUDE_SCHEMA_OPTION = typer.Option( + [], + "--exclude-schema", + help="Skip files referencing this DB schema (repeatable); " + "files with no detectable schema reference are never excluded", +) -def _as_area_set(values: list[str]) -> frozenset[str] | None: +def _as_set(values: list[str]) -> frozenset[str] | None: return frozenset(values) if values else None testdb_app = typer.Typer(name="testdb", help="Manage the shared local Postgres test container") @@ -59,17 +71,20 @@ def compare( dialect: str = typer.Option("postgres", "--dialect", help="postgres (default) or mssql"), area: list[str] = _AREA_OPTION, exclude_area: list[str] = _EXCLUDE_AREA_OPTION, + schema: list[str] = _SCHEMA_OPTION, + exclude_schema: list[str] = _EXCLUDE_SCHEMA_OPTION, scripts_dir: Path = typer.Argument(..., help="Directory containing SQL scripts"), ) -> None: """Compare SQL scripts to a live database and report differences.""" if not scripts_dir.is_dir(): err_console.print(f"[red]Error:[/red] {scripts_dir} is not a directory") raise typer.Exit(2) - if report_extra_db and (area or exclude_area): + if report_extra_db and (area or exclude_area or schema or exclude_schema): console.print( - "[yellow]⚠[/yellow] --report-extra-db with --area/--exclude-area will report every DB object " - "outside the filtered area(s) as \"missing in scripts\", since the live database has no concept " - "of areas — only the scripts side is filtered." + "[yellow]⚠[/yellow] --report-extra-db with --area/--exclude-area/--schema/--exclude-schema will " + "report every DB object outside the filtered area(s)/schema(s) as \"missing in scripts\", since the " + "live database has no concept of areas — and isn't itself filtered by --schema either — only the " + "scripts side is filtered." ) try: @@ -91,7 +106,12 @@ def compare( with console.status("Parsing SQL scripts..."): scripts_schema = parse_directory( - scripts_dir, dialect=backend.dialect, areas=_as_area_set(area), exclude_areas=_as_area_set(exclude_area) + scripts_dir, + dialect=backend.dialect, + areas=_as_set(area), + exclude_areas=_as_set(exclude_area), + schemas=_as_set(schema), + exclude_schemas=_as_set(exclude_schema), ) with console.status("Introspecting database..."): @@ -188,17 +208,33 @@ def fetch_missing( @testdb_app.command("up") -def testdb_up() -> None: +def testdb_up( + area: list[str] = _AREA_OPTION, + exclude_area: list[str] = _EXCLUDE_AREA_OPTION, + schema: list[str] = _SCHEMA_OPTION, + exclude_schema: list[str] = _EXCLUDE_SCHEMA_OPTION, +) -> None: """Ensure the container is running, the workspace DB exists, and schema is applied.""" - testdb.ensure_testdb() + testdb.ensure_testdb( + areas=_as_set(area), exclude_areas=_as_set(exclude_area), schemas=_as_set(schema), + exclude_schemas=_as_set(exclude_schema), + ) 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( + area: list[str] = _AREA_OPTION, + exclude_area: list[str] = _EXCLUDE_AREA_OPTION, + schema: list[str] = _SCHEMA_OPTION, + exclude_schema: list[str] = _EXCLUDE_SCHEMA_OPTION, +) -> None: """Drop and recreate only this workspace's database, then reapply schema + seed data.""" - testdb.reset_testdb() + testdb.reset_testdb( + areas=_as_set(area), exclude_areas=_as_set(exclude_area), schemas=_as_set(schema), + exclude_schemas=_as_set(exclude_schema), + ) info = testdb.status() console.print(f"[green]Test DB reset:[/green] {info['database']}") @@ -272,6 +308,8 @@ def migrate_check( ), area: list[str] = _AREA_OPTION, exclude_area: list[str] = _EXCLUDE_AREA_OPTION, + schema: list[str] = _SCHEMA_OPTION, + exclude_schema: list[str] = _EXCLUDE_SCHEMA_OPTION, ) -> None: """List which migration files under migrations_dir are applied vs. pending.""" if not migrations_dir.is_dir(): @@ -281,7 +319,11 @@ 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_set(area), + exclude_areas=_as_set(exclude_area), + schemas=_as_set(schema), + exclude_schemas=_as_set(exclude_schema), ) try: applied = migrate.applied_migrations(conninfo, tracking_table) @@ -323,6 +365,8 @@ 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, + schema: list[str] = _SCHEMA_OPTION, + exclude_schema: list[str] = _EXCLUDE_SCHEMA_OPTION, ) -> None: """Apply pending migration files, in filename order, tracking each in tracking_table.""" if not migrations_dir.is_dir(): @@ -331,7 +375,8 @@ def migrate_apply( conninfo = build_conninfo(url, entra_user) tracking_table = tracking_table or migrate.default_tracking_table(migrations_dir) - areas, exclude_areas = _as_area_set(area), _as_area_set(exclude_area) + areas, exclude_areas = _as_set(area), _as_set(exclude_area) + schemas, exclude_schemas = _as_set(schema), _as_set(exclude_schema) target_desc = url.rsplit("@", 1)[-1] if "@" in url else url if not yes: typer.confirm(f"About to run migrations against {target_desc}. Continue?", abort=True) @@ -341,13 +386,22 @@ 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, + schemas=schemas, + exclude_schemas=exclude_schemas, ) 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, schemas=schemas, + exclude_schemas=exclude_schemas, + ) if not targets: console.print("No pending migrations.") diff --git a/pgdevkit/migrate.py b/pgdevkit/migrate.py index ad0289e..6d5eae3 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 .schemas import filter_by_schema _IDENTIFIER = r"[A-Za-z_][A-Za-z0-9_]*" _DEFAULT_TRACKING_TABLE = "public.schema_migrations" @@ -281,9 +282,12 @@ def list_migration_files( *, areas: frozenset[str] | None = None, exclude_areas: frozenset[str] | None = None, + schemas: frozenset[str] | None = None, + exclude_schemas: frozenset[str] | None = None, ) -> list[Path]: 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 filter_by_schema(files, only=schemas, exclude=exclude_schemas) def applied_migrations(conninfo: str, tracking_table: str) -> dict[str, tuple[datetime, str]]: @@ -306,9 +310,13 @@ def pending_migrations( *, areas: frozenset[str] | None = None, exclude_areas: frozenset[str] | None = None, + schemas: frozenset[str] | None = None, + exclude_schemas: frozenset[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, schemas=schemas, exclude_schemas=exclude_schemas + ) return [p for p in files if p.name not in applied] diff --git a/pgdevkit/parser.py b/pgdevkit/parser.py index c89c6d5..7953604 100644 --- a/pgdevkit/parser.py +++ b/pgdevkit/parser.py @@ -10,6 +10,7 @@ from .areas import area_allowed, parse_areas from .dialect import Dialect, POSTGRES, resolve_dialect +from .schemas import schema_allowed, sql_schemas from .models import ( ColumnDef, ConstraintDef, CompositeTypeDef, DatabaseSchema, EnumDef, FunctionDef, IndexDef, TableDef, ViewDef, @@ -55,15 +56,21 @@ def parse_directory( dialect: str | Dialect = "postgres", areas: frozenset[str] | None = None, exclude_areas: frozenset[str] | None = None, + schemas: frozenset[str] | None = None, + exclude_schemas: frozenset[str] | None = None, ) -> DatabaseSchema: resolved = resolve_dialect(dialect) db_schema = DatabaseSchema() for sql_file in sorted(_iter_sql_files(scripts_dir)): - # Read once and reuse for both the area check and parsing, rather than - # filtering the file list up front (which would need its own read). + # Read once and reuse for the area/schema checks and parsing, rather + # than filtering the file list up front (which would need its own read). content = sql_file.read_text(encoding="utf-8") if (areas or exclude_areas) and not area_allowed(parse_areas(content), only=areas, exclude=exclude_areas): continue + if (schemas or exclude_schemas) and not schema_allowed( + sql_schemas(content, resolved), only=schemas, exclude=exclude_schemas + ): + continue _parse_file(sql_file, db_schema, resolved, content=content) return db_schema diff --git a/pgdevkit/schemas.py b/pgdevkit/schemas.py new file mode 100644 index 0000000..cdf41aa --- /dev/null +++ b/pgdevkit/schemas.py @@ -0,0 +1,112 @@ +"""Schema-membership filtering, composable with (but independent of) the +`-- area:` tag filtering in `areas.py`. + +Unlike area, which is an explicit opt-in tag, a file's schema membership is +derived by parsing its SQL: every schema-qualified (or default-schema, when +unqualified) table/view/function/index/schema reference across every +statement in the file, DDL or DML alike. + +A file whose schema(s) can't be determined -- content sqlglot can't parse at +all, or with no table/schema reference in it (e.g. a DO block touching no +table) -- is treated the same as an untagged file for `-- area:`: `only` +filters always keep it, `exclude` filters never drop it. Being unable to +prove a file belongs to an excluded schema is not the same as proving it +doesn't, so this errs toward keeping the file in scope rather than silently +dropping it. +""" + +from __future__ import annotations + +import re +from pathlib import Path + +import sqlglot +import sqlglot.expressions as exp + +from .dialect import Dialect, POSTGRES + +# Schemas that hold system catalog views/tables, never a file this project +# manages -- a reference to one (e.g. an idempotency guard querying +# information_schema/pg_catalog/sys) is never a real schema membership signal. +_SYSTEM_SCHEMAS = {"pg_catalog", "information_schema", "sys"} + +_CREATE_SCHEMA_RE = re.compile( + r"CREATE\s+SCHEMA\s+(?:IF\s+NOT\s+EXISTS\s+)?(?:AUTHORIZATION\s+)?\"?(\w+)\"?", re.IGNORECASE +) +_QUALIFIED_REF_RE = re.compile(r"\b(\w+)\.\w+") + + +def _regex_fallback(sql: str) -> frozenset[str]: + """Crude schema scan used when sqlglot can't parse `sql` at all. Errs + toward over-matching (any `schema.name`-shaped token) rather than + under-matching -- a false positive here only keeps a file in scope for + one extra schema filter, where a false negative would silently drop it + from an `only` filter.""" + schemas = set(_CREATE_SCHEMA_RE.findall(sql)) + schemas.update(m.group(1) for m in _QUALIFIED_REF_RE.finditer(sql)) + return frozenset(s for s in schemas if s.lower() not in _SYSTEM_SCHEMAS) + + +def sql_schemas(content: str, dialect: Dialect = POSTGRES) -> frozenset[str]: + """Schema names referenced anywhere in `content`: every table/view/ + function/index reference's schema (its `dialect.default_schema` when + unqualified), plus any `CREATE SCHEMA name`. Falls back to a regex scan + when sqlglot can't parse the content, or finds no reference at all.""" + try: + exprs = sqlglot.parse(content, dialect=dialect.sqlglot_name, error_level=sqlglot.ErrorLevel.IGNORE) + except Exception: # noqa: BLE001 + return _regex_fallback(content) + + schemas: set[str] = set() + for e in exprs: + if e is None: + continue + for t in e.find_all(exp.Table): + # T-SQL's `EXEC('...dynamic sql...')` parses as a Table subquery + # whose "name" is the whole literal string, not a real + # identifier -- e.g. a quoted "CREATE SCHEMA app". Guard against + # treating that as a schema-qualified (or default-schema) table + # reference by requiring the name to actually look like one. + if not re.fullmatch(r"\w+", t.name or ""): + continue + db_node = t.args.get("db") + name = db_node.name if db_node else dialect.default_schema + if name.lower() not in _SYSTEM_SCHEMAS: + schemas.add(name) + + if not schemas: + return _regex_fallback(content) + return frozenset(schemas) + + +def file_schemas(path: Path, dialect: Dialect = POSTGRES) -> frozenset[str]: + """Schema names referenced in the file at `path`.""" + return sql_schemas(path.read_text(encoding="utf-8"), dialect) + + +def schema_allowed( + schemas: frozenset[str], + *, + only: frozenset[str] | None = None, + exclude: frozenset[str] | None = None, +) -> bool: + """Whether a file that references `schemas` passes an `only`/`exclude` filter.""" + if exclude and schemas & exclude: + return False + if only and schemas and not (schemas & only): + return False + return True + + +def filter_by_schema( + paths: list[Path], + *, + only: frozenset[str] | None = None, + exclude: frozenset[str] | None = None, + dialect: Dialect = POSTGRES, +) -> list[Path]: + """`paths` restricted by an `only`/`exclude` schema filter. Returns `paths` + unchanged (no file reads) when neither filter is set.""" + if not only and not exclude: + return paths + return [p for p in paths if schema_allowed(file_schemas(p, dialect), only=only, exclude=exclude)] diff --git a/pgdevkit/testdb/api.py b/pgdevkit/testdb/api.py index 2e0e85a..48424e2 100644 --- a/pgdevkit/testdb/api.py +++ b/pgdevkit/testdb/api.py @@ -68,24 +68,53 @@ 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, + *, + areas: frozenset[str] | None = None, + exclude_areas: frozenset[str] | None = None, + schemas: frozenset[str] | None = None, + exclude_schemas: frozenset[str] | None = None, +) -> 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, + areas=areas, + exclude_areas=exclude_areas, + schemas=schemas, + exclude_schemas=exclude_schemas, ) -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, + *, + areas: frozenset[str] | None = None, + exclude_areas: frozenset[str] | None = None, + schemas: frozenset[str] | None = None, + exclude_schemas: frozenset[str] | None = None, +) -> 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`). + + `areas`/`exclude_areas` and `schemas`/`exclude_schemas` restrict which + database/ files get applied -- e.g. for a test DB scoped to one area or + schema. Neither filters what gets *dropped* by force_reset/clean, only + what gets (re)applied.""" 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, + areas=areas, exclude_areas=exclude_areas, schemas=schemas, exclude_schemas=exclude_schemas, + ) ensure_container() @@ -93,16 +122,29 @@ 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, + areas=areas, exclude_areas=exclude_areas, schemas=schemas, exclude_schemas=exclude_schemas, + ) 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, + *, + areas: frozenset[str] | None = None, + exclude_areas: frozenset[str] | None = None, + schemas: frozenset[str] | None = None, + exclude_schemas: frozenset[str] | None = None, +) -> 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, + areas=areas, exclude_areas=exclude_areas, schemas=schemas, exclude_schemas=exclude_schemas, + ) 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..b310d60 100644 --- a/pgdevkit/testdb/mssql/api.py +++ b/pgdevkit/testdb/mssql/api.py @@ -101,7 +101,16 @@ 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, + *, + areas: frozenset[str] | None = None, + exclude_areas: frozenset[str] | None = None, + schemas: frozenset[str] | None = None, + exclude_schemas: frozenset[str] | None = None, +) -> None: def _connect() -> Any: return mssql_python.connect(_db_dsn(db_name), autocommit=True) @@ -110,7 +119,10 @@ 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, + areas=areas, exclude_areas=exclude_areas, schemas=schemas, exclude_schemas=exclude_schemas, + ): for batch in query.split_tsql_batches(sql): def _exec(batch: str = batch) -> None: @@ -126,14 +138,26 @@ def _exec(batch: str = batch) -> None: 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, + *, + areas: frozenset[str] | None = None, + exclude_areas: frozenset[str] | None = None, + schemas: frozenset[str] | None = None, + exclude_schemas: frozenset[str] | None = None, +) -> 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, + areas=areas, exclude_areas=exclude_areas, schemas=schemas, exclude_schemas=exclude_schemas, + ) asyncio.run(_run()) return _env_for(config, db_name) diff --git a/pgdevkit/testdb/schema.py b/pgdevkit/testdb/schema.py index 581e08b..10f940f 100644 --- a/pgdevkit/testdb/schema.py +++ b/pgdevkit/testdb/schema.py @@ -13,9 +13,11 @@ from psycopg.rows import dict_row from psycopg.sql import SQL, Identifier, Placeholder +from ..areas import area_allowed, parse_areas from ..db.complex_types import ComplexHelper from ..dialect import Dialect, POSTGRES from ..parser import IGNORED_DIR_NAMES +from ..schemas import schema_allowed, sql_schemas logger = logging.getLogger(__name__) logging.getLogger("sqlglot").setLevel(logging.ERROR) @@ -119,8 +121,21 @@ def _get_sql_deps(sql: str, dialect: Dialect = POSTGRES) -> set[str]: return deps -def _iter_sql_files(database_dir: Path, dialect: Dialect = POSTGRES): - """Yield (Path, sql_content) pairs in dependency-safe execution order.""" +def _iter_sql_files( + database_dir: Path, + dialect: Dialect = POSTGRES, + *, + areas: frozenset[str] | None = None, + exclude_areas: frozenset[str] | None = None, + schemas: frozenset[str] | None = None, + exclude_schemas: frozenset[str] | None = None, +): + """Yield (Path, sql_content) pairs in dependency-safe execution order. + + A file dropped by the area/schema filter is skipped entirely -- as if it + didn't exist -- so it never delivers a dependency another (in-scope) + file waits on; that's the same tradeoff `list_migration_files`/ + `parse_directory` make for migrations/database code files.""" files: list[Path] = [] for root, dirs, dbfiles in os.walk(database_dir): dirs[:] = [d for d in dirs if d not in IGNORED_DIR_NAMES] @@ -138,6 +153,12 @@ def _iter_sql_files(database_dir: Path, dialect: Dialect = POSTGRES): for file in sorted(files, key=lambda p: (_get_type_order(p), p.name)): content = file.read_text(encoding="utf-8") + if (areas or exclude_areas) and not area_allowed(parse_areas(content), only=areas, exclude=exclude_areas): + continue + if (schemas or exclude_schemas) and not schema_allowed( + sql_schemas(content, dialect), only=schemas, exclude=exclude_schemas + ): + continue deps = _get_sql_deps(content, dialect) if file.parent.name in _SCHEMA_QUALIFIED_TYPES: schema = _strip_layer_prefix(file.parent.parent.name) @@ -242,6 +263,10 @@ async def apply_schema( force_reset: bool = False, *, dialect: Dialect = POSTGRES, + areas: frozenset[str] | None = None, + exclude_areas: frozenset[str] | None = None, + schemas: frozenset[str] | None = None, + exclude_schemas: frozenset[str] | None = None, ) -> None: """Apply every .sql file under database_dir (in dependency-safe order) and seed any matching .test_data.json files. Safe to call repeatedly. @@ -249,7 +274,12 @@ async def apply_schema( `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 - `database_dir` must reflect the current, final schema on their own.""" + `database_dir` must reflect the current, final schema on their own. + + `areas`/`exclude_areas` and `schemas`/`exclude_schemas` restrict which + files get applied, the same as `pgdevkit migrate`/`compare` — handy for + standing up a test DB scoped to one area or schema instead of the whole + project.""" await con.set_autocommit(True) for extension in extensions: await con.execute(SQL("CREATE EXTENSION IF NOT EXISTS {e}").format(e=Identifier(extension))) @@ -265,7 +295,14 @@ async def _apply(file: Path, sql: str) -> None: 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, + areas=areas, + exclude_areas=exclude_areas, + schemas=schemas, + exclude_schemas=exclude_schemas, + ): try: await _apply(file, sql) except Exception as e: # noqa: BLE001 diff --git a/tests/test_migrate_schemas.py b/tests/test_migrate_schemas.py new file mode 100644 index 0000000..b0d5d8b --- /dev/null +++ b/tests/test_migrate_schemas.py @@ -0,0 +1,52 @@ +from __future__ import annotations + +from pathlib import Path + +from pgdevkit.migrate import list_migration_files + + +def _write(dir: Path, name: str, content: str) -> Path: + p = dir / name + p.write_text(content, encoding="utf-8") + return p + + +class TestListMigrationFilesSchemaFiltering: + def test_no_filters_lists_everything(self, tmp_path: Path): + a = _write(tmp_path, "001_a.sql", "CREATE TABLE billing.a (id int);\n") + b = _write(tmp_path, "002_b.sql", "CREATE TABLE reporting.b (id int);\n") + assert list_migration_files(tmp_path) == sorted([a, b]) + + def test_schemas_filter_keeps_matching(self, tmp_path: Path): + billing = _write(tmp_path, "001_billing.sql", "CREATE TABLE billing.a (id int);\n") + reporting = _write(tmp_path, "002_reporting.sql", "CREATE TABLE reporting.b (id int);\n") + + result = list_migration_files(tmp_path, schemas=frozenset({"billing"})) + assert result == [billing] + assert reporting not in result + + def test_exclude_schemas_drops_matching(self, tmp_path: Path): + billing = _write(tmp_path, "001_billing.sql", "CREATE TABLE billing.a (id int);\n") + reporting = _write(tmp_path, "002_reporting.sql", "CREATE TABLE reporting.b (id int);\n") + + result = list_migration_files(tmp_path, exclude_schemas=frozenset({"billing"})) + assert result == [reporting] + assert billing not in result + + def test_undetectable_schema_always_kept(self, tmp_path: Path): + undetectable = _write(tmp_path, "001_common.sql", "select 1;\n") + + assert list_migration_files(tmp_path, schemas=frozenset({"billing"})) == [undetectable] + assert list_migration_files(tmp_path, exclude_schemas=frozenset({"billing"})) == [undetectable] + + def test_area_and_schema_filters_combine(self, tmp_path: Path): + keep = _write(tmp_path, "001_keep.sql", "-- area: billing\nCREATE TABLE reporting.a (id int);\n") + wrong_area = _write(tmp_path, "002_wrong_area.sql", "-- area: ops\nCREATE TABLE reporting.b (id int);\n") + wrong_schema = _write( + tmp_path, "003_wrong_schema.sql", "-- area: billing\nCREATE TABLE billing.c (id int);\n" + ) + + result = list_migration_files(tmp_path, areas=frozenset({"billing"}), schemas=frozenset({"reporting"})) + assert result == [keep] + assert wrong_area not in result + assert wrong_schema not in result diff --git a/tests/test_parser_schemas.py b/tests/test_parser_schemas.py new file mode 100644 index 0000000..fe96b56 --- /dev/null +++ b/tests/test_parser_schemas.py @@ -0,0 +1,44 @@ +from __future__ import annotations + +from pathlib import Path + +from pgdevkit.parser import parse_directory + + +def _write(dir: Path, name: str, content: str) -> Path: + p = dir / name + p.write_text(content, encoding="utf-8") + return p + + +class TestParseDirectorySchemaFiltering: + def test_no_filters_parses_everything(self, tmp_path: Path): + _write(tmp_path, "billing.sql", "CREATE TABLE billing.a (id int);\n") + _write(tmp_path, "reporting.sql", "CREATE TABLE reporting.b (id int);\n") + schema = parse_directory(tmp_path) + assert set(schema.tables) == {"billing.a", "reporting.b"} + + def test_schemas_filter_keeps_matching_and_undetectable(self, tmp_path: Path): + _write(tmp_path, "billing.sql", "CREATE TABLE billing.a (id int);\n") + _write(tmp_path, "reporting.sql", "CREATE TABLE reporting.b (id int);\n") + + schema = parse_directory(tmp_path, schemas=frozenset({"billing"})) + assert set(schema.tables) == {"billing.a"} + + def test_exclude_schemas_drops_matching(self, tmp_path: Path): + _write(tmp_path, "billing.sql", "CREATE TABLE billing.a (id int);\n") + _write(tmp_path, "reporting.sql", "CREATE TABLE reporting.b (id int);\n") + + schema = parse_directory(tmp_path, exclude_schemas=frozenset({"billing"})) + assert set(schema.tables) == {"reporting.b"} + + def test_area_and_schema_filters_combine(self, tmp_path: Path): + # Passes both filters. + _write(tmp_path, "keep.sql", "-- area: billing\nCREATE TABLE reporting.a (id int);\n") + # Wrong area. + _write(tmp_path, "wrong_area.sql", "-- area: ops\nCREATE TABLE reporting.b (id int);\n") + # Wrong schema. + _write(tmp_path, "wrong_schema.sql", "-- area: billing\nCREATE TABLE billing.c (id int);\n") + + schema = parse_directory(tmp_path, areas=frozenset({"billing"}), schemas=frozenset({"reporting"})) + assert set(schema.tables) == {"reporting.a"} diff --git a/tests/test_schemas.py b/tests/test_schemas.py new file mode 100644 index 0000000..d80476f --- /dev/null +++ b/tests/test_schemas.py @@ -0,0 +1,130 @@ +from __future__ import annotations + +from pathlib import Path + +from pgdevkit.dialect import MSSQL, POSTGRES +from pgdevkit.schemas import file_schemas, filter_by_schema, schema_allowed, sql_schemas + + +class TestSqlSchemas: + def test_unqualified_create_table_uses_default_schema(self): + assert sql_schemas("CREATE TABLE t (id int);\n") == frozenset({"public"}) + + def test_qualified_create_table(self): + assert sql_schemas("CREATE TABLE billing.invoice (id int);\n") == frozenset({"billing"}) + + def test_multiple_schemas_referenced(self): + sql = "CREATE TABLE billing.invoice (id int, customer_id int references reporting.customer(id));\n" + assert sql_schemas(sql) == frozenset({"billing", "reporting"}) + + def test_create_schema_statement(self): + assert sql_schemas("CREATE SCHEMA billing;\n") == frozenset({"billing"}) + + def test_alter_table(self): + assert sql_schemas("ALTER TABLE billing.invoice ADD COLUMN paid boolean;\n") == frozenset({"billing"}) + + def test_drop_table(self): + assert sql_schemas("DROP TABLE billing.invoice;\n") == frozenset({"billing"}) + + def test_data_migration_statements(self): + assert sql_schemas("INSERT INTO billing.invoice (id) VALUES (1);\n") == frozenset({"billing"}) + assert sql_schemas("UPDATE billing.invoice SET paid = true;\n") == frozenset({"billing"}) + assert sql_schemas("DELETE FROM billing.invoice WHERE id = 1;\n") == frozenset({"billing"}) + + def test_system_schema_references_excluded(self): + sql = "SELECT 1 FROM pg_catalog.pg_class WHERE relname = 'invoice'" + assert sql_schemas(sql) == frozenset() + + def test_information_schema_reference_excluded(self): + sql = "SELECT 1 FROM information_schema.columns WHERE table_name = 'invoice'" + assert sql_schemas(sql) == frozenset() + + def test_no_table_reference_is_undetectable(self): + assert sql_schemas("SELECT 1;\n") == frozenset() + + def test_mssql_default_schema(self): + assert sql_schemas("CREATE TABLE t (id int);\n", MSSQL) == frozenset({"dbo"}) + + def test_mssql_sys_schema_excluded(self): + # The idempotency-guard SELECT against sys.schemas must not itself + # count as a schema reference; the dynamic `EXEC('CREATE SCHEMA + # app')` this guards is still a real (if indirect) declaration of + # "app" -- picked up by the regex fallback once sqlglot's own + # mis-parse of the EXEC(...) literal as a table name is discarded. + sql = "IF NOT EXISTS (SELECT 1 FROM sys.schemas WHERE name = 'app') BEGIN EXEC('CREATE SCHEMA app'); END" + assert sql_schemas(sql, MSSQL) == frozenset({"app"}) + + def test_unparseable_content_falls_back_to_regex(self): + # Deliberately malformed SQL that still contains a schema-qualified + # reference sqlglot can't make sense of as a whole statement. + sql = "!!! not sql billing.invoice !!!" + assert sql_schemas(sql) == frozenset({"billing"}) + + +class TestFileSchemas: + def test_reads_from_disk(self, tmp_path: Path): + f = tmp_path / "001_thing.sql" + f.write_text("CREATE TABLE billing.invoice (id int);\n", encoding="utf-8") + assert file_schemas(f) == frozenset({"billing"}) + + +class TestSchemaAllowed: + def test_no_filters_always_allowed(self): + assert schema_allowed(frozenset({"billing"})) is True + assert schema_allowed(frozenset()) is True + + def test_only_filter_undetectable_always_passes(self): + assert schema_allowed(frozenset(), only=frozenset({"billing"})) is True + + def test_only_filter_matching_schema_passes(self): + assert schema_allowed(frozenset({"billing"}), only=frozenset({"billing"})) is True + + def test_only_filter_non_matching_schema_fails(self): + assert schema_allowed(frozenset({"reporting"}), only=frozenset({"billing"})) is False + + def test_exclude_filter_undetectable_never_dropped(self): + assert schema_allowed(frozenset(), exclude=frozenset({"billing"})) is True + + def test_exclude_filter_matching_schema_dropped(self): + assert schema_allowed(frozenset({"billing"}), exclude=frozenset({"billing"})) is False + + def test_exclude_filter_non_matching_schema_passes(self): + assert schema_allowed(frozenset({"reporting"}), exclude=frozenset({"billing"})) is True + + def test_only_and_exclude_combined_exclude_wins(self): + schemas = frozenset({"billing"}) + assert schema_allowed(schemas, only=frozenset({"billing"}), exclude=frozenset({"billing"})) is False + + +class TestFilterBySchema: + def test_no_filters_returns_paths_unchanged(self, tmp_path: Path): + paths = [tmp_path / "a.sql", tmp_path / "b.sql"] + assert filter_by_schema(paths) == paths + + def test_only_keeps_matching_and_undetectable(self, tmp_path: Path): + billing = tmp_path / "billing.sql" + billing.write_text("CREATE TABLE billing.invoice (id int);\n", encoding="utf-8") + reporting = tmp_path / "reporting.sql" + reporting.write_text("CREATE TABLE reporting.customer (id int);\n", encoding="utf-8") + undetectable = tmp_path / "common.sql" + undetectable.write_text("SELECT 1;\n", encoding="utf-8") + + result = filter_by_schema([billing, reporting, undetectable], only=frozenset({"billing"})) + assert set(result) == {billing, undetectable} + + def test_exclude_drops_matching_but_keeps_undetectable(self, tmp_path: Path): + billing = tmp_path / "billing.sql" + billing.write_text("CREATE TABLE billing.invoice (id int);\n", encoding="utf-8") + undetectable = tmp_path / "common.sql" + undetectable.write_text("SELECT 1;\n", encoding="utf-8") + + result = filter_by_schema([billing, undetectable], exclude=frozenset({"billing"})) + assert result == [undetectable] + + def test_dialect_affects_default_schema(self, tmp_path: Path): + f = tmp_path / "thing.sql" + f.write_text("CREATE TABLE t (id int);\n", encoding="utf-8") + + assert filter_by_schema([f], only=frozenset({"public"}), dialect=POSTGRES) == [f] + assert filter_by_schema([f], only=frozenset({"public"}), dialect=MSSQL) == [] + assert filter_by_schema([f], only=frozenset({"dbo"}), dialect=MSSQL) == [f] diff --git a/tests/testdb/test_schema_filtering.py b/tests/testdb/test_schema_filtering.py new file mode 100644 index 0000000..c8a4354 --- /dev/null +++ b/tests/testdb/test_schema_filtering.py @@ -0,0 +1,62 @@ +from __future__ import annotations + +from pathlib import Path + +from pgdevkit.testdb.schema import _iter_sql_files + + +def _write(base: Path, rel: str, content: str) -> Path: + p = base / rel + p.parent.mkdir(parents=True, exist_ok=True) + p.write_text(content, encoding="utf-8") + return p + + +def _names(base: Path, **kwargs) -> set[Path]: + return {f.relative_to(base) for f, _ in _iter_sql_files(base, **kwargs)} + + +class TestIterSqlFilesAreaFiltering: + def test_no_filters_yields_everything(self, tmp_path: Path): + _write(tmp_path, "billing/tables/x.sql", "-- area: billing\nCREATE TABLE billing.x (id int);\n") + _write(tmp_path, "reporting/tables/y.sql", "-- area: reporting\nCREATE TABLE reporting.y (id int);\n") + assert _names(tmp_path) == {Path("billing/tables/x.sql"), Path("reporting/tables/y.sql")} + + def test_areas_filter_keeps_matching_and_untagged(self, tmp_path: Path): + _write(tmp_path, "billing/tables/x.sql", "-- area: billing\nCREATE TABLE billing.x (id int);\n") + _write(tmp_path, "reporting/tables/y.sql", "-- area: reporting\nCREATE TABLE reporting.y (id int);\n") + _write(tmp_path, "common/tables/z.sql", "CREATE TABLE common.z (id int);\n") + + result = _names(tmp_path, areas=frozenset({"billing"})) + assert result == {Path("billing/tables/x.sql"), Path("common/tables/z.sql")} + + def test_exclude_areas_drops_matching_but_keeps_untagged(self, tmp_path: Path): + _write(tmp_path, "billing/tables/x.sql", "-- area: billing\nCREATE TABLE billing.x (id int);\n") + _write(tmp_path, "common/tables/z.sql", "CREATE TABLE common.z (id int);\n") + + result = _names(tmp_path, exclude_areas=frozenset({"billing"})) + assert result == {Path("common/tables/z.sql")} + + +class TestIterSqlFilesSchemaFiltering: + def test_schemas_filter_keeps_matching(self, tmp_path: Path): + _write(tmp_path, "billing/tables/x.sql", "CREATE TABLE billing.x (id int);\n") + _write(tmp_path, "reporting/tables/y.sql", "CREATE TABLE reporting.y (id int);\n") + + result = _names(tmp_path, schemas=frozenset({"billing"})) + assert result == {Path("billing/tables/x.sql")} + + def test_exclude_schemas_drops_matching(self, tmp_path: Path): + _write(tmp_path, "billing/tables/x.sql", "CREATE TABLE billing.x (id int);\n") + _write(tmp_path, "reporting/tables/y.sql", "CREATE TABLE reporting.y (id int);\n") + + result = _names(tmp_path, exclude_schemas=frozenset({"billing"})) + assert result == {Path("reporting/tables/y.sql")} + + def test_area_and_schema_filters_combine(self, tmp_path: Path): + _write(tmp_path, "keep/tables/a.sql", "-- area: billing\nCREATE TABLE reporting.a (id int);\n") + _write(tmp_path, "wrong_area/tables/b.sql", "-- area: ops\nCREATE TABLE reporting.b (id int);\n") + _write(tmp_path, "wrong_schema/tables/c.sql", "-- area: billing\nCREATE TABLE billing.c (id int);\n") + + result = _names(tmp_path, areas=frozenset({"billing"}), schemas=frozenset({"reporting"})) + assert result == {Path("keep/tables/a.sql")} From 97c8e0094a524b2bda1fb003c3a4c475e9e88768 Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Mon, 7 Sep 2026 20:36:44 +0200 Subject: [PATCH 2/5] Fix schema detection dropping bare CREATE SCHEMA when file has other refs The EXEC(...)-literal-misparse guard in sql_schemas() rejected any Table node without an identifier-shaped name, including the legitimate empty name sqlglot gives a bare CREATE SCHEMA (whose name lives in `db`, not `this`). That silently dropped the CREATE SCHEMA's schema whenever the same file had another real table reference too (the empty-schemas regex-fallback safety net never triggered). Found by code review. Also de-duplicates the system-schema exclusion list (pg_catalog/ information_schema/sys) into dialect.SYSTEM_SCHEMAS, shared by schemas.py and testdb/schema.py instead of each keeping its own copy. Co-Authored-By: Claude Sonnet 5 Claude-Session: https://claude.ai/code/session_01StWgZw5TDd7KNGY3VzNGMg --- pgdevkit/dialect.py | 10 ++++++++++ pgdevkit/schemas.py | 18 ++++++++---------- pgdevkit/testdb/schema.py | 24 +++++++++++------------- tests/test_schemas.py | 8 ++++++++ 4 files changed, 37 insertions(+), 23 deletions(-) diff --git a/pgdevkit/dialect.py b/pgdevkit/dialect.py index b414695..27e16b3 100644 --- a/pgdevkit/dialect.py +++ b/pgdevkit/dialect.py @@ -43,6 +43,16 @@ } +# Schemas that hold system catalog views/tables, never a file any project +# using pgdevkit manages -- a reference to one (e.g. an idempotency guard +# querying it, or a `SELECT ... FROM information_schema/pg_catalog/sys ...`) +# is never a real schema-membership or cross-file-dependency signal. Shared +# by `schemas.py` (schema-reference filtering) and `testdb/schema.py` +# (dependency-safe apply ordering), which both walk the same sqlglot Table +# nodes for a related-but-different purpose. +SYSTEM_SCHEMAS = {"pg_catalog", "information_schema", "sys"} + + @dataclass(frozen=True) class Dialect: """A thin wrapper around a sqlglot dialect name plus the handful of diff --git a/pgdevkit/schemas.py b/pgdevkit/schemas.py index cdf41aa..8224792 100644 --- a/pgdevkit/schemas.py +++ b/pgdevkit/schemas.py @@ -23,12 +23,7 @@ import sqlglot import sqlglot.expressions as exp -from .dialect import Dialect, POSTGRES - -# Schemas that hold system catalog views/tables, never a file this project -# manages -- a reference to one (e.g. an idempotency guard querying -# information_schema/pg_catalog/sys) is never a real schema membership signal. -_SYSTEM_SCHEMAS = {"pg_catalog", "information_schema", "sys"} +from .dialect import Dialect, POSTGRES, SYSTEM_SCHEMAS _CREATE_SCHEMA_RE = re.compile( r"CREATE\s+SCHEMA\s+(?:IF\s+NOT\s+EXISTS\s+)?(?:AUTHORIZATION\s+)?\"?(\w+)\"?", re.IGNORECASE @@ -44,7 +39,7 @@ def _regex_fallback(sql: str) -> frozenset[str]: from an `only` filter.""" schemas = set(_CREATE_SCHEMA_RE.findall(sql)) schemas.update(m.group(1) for m in _QUALIFIED_REF_RE.finditer(sql)) - return frozenset(s for s in schemas if s.lower() not in _SYSTEM_SCHEMAS) + return frozenset(s for s in schemas if s.lower() not in SYSTEM_SCHEMAS) def sql_schemas(content: str, dialect: Dialect = POSTGRES) -> frozenset[str]: @@ -66,12 +61,15 @@ def sql_schemas(content: str, dialect: Dialect = POSTGRES) -> frozenset[str]: # whose "name" is the whole literal string, not a real # identifier -- e.g. a quoted "CREATE SCHEMA app". Guard against # treating that as a schema-qualified (or default-schema) table - # reference by requiring the name to actually look like one. - if not re.fullmatch(r"\w+", t.name or ""): + # reference by requiring a *non-empty* name to actually look like + # one -- but still accept an empty name, which is the legitimate + # shape sqlglot gives a bare `CREATE SCHEMA x` (whose schema + # name ends up in `db`, not `this`). + if t.name and not re.fullmatch(r"\w+", t.name): continue db_node = t.args.get("db") name = db_node.name if db_node else dialect.default_schema - if name.lower() not in _SYSTEM_SCHEMAS: + if name.lower() not in SYSTEM_SCHEMAS: schemas.add(name) if not schemas: diff --git a/pgdevkit/testdb/schema.py b/pgdevkit/testdb/schema.py index 10f940f..a4e8f49 100644 --- a/pgdevkit/testdb/schema.py +++ b/pgdevkit/testdb/schema.py @@ -15,7 +15,7 @@ from ..areas import area_allowed, parse_areas from ..db.complex_types import ComplexHelper -from ..dialect import Dialect, POSTGRES +from ..dialect import Dialect, POSTGRES, SYSTEM_SCHEMAS from ..parser import IGNORED_DIR_NAMES from ..schemas import schema_allowed, sql_schemas @@ -67,16 +67,14 @@ def _strip_layer_prefix(schema_name: str) -> str: } -# Schemas that hold system catalog views/tables, never a file this project -# manages -- a reference to one is never a real cross-file dependency to -# wait for. Matters most for T-SQL, where "IF NOT EXISTS (SELECT ... FROM -# sys.schemas/sys.tables/sys.objects ...) BEGIN CREATE ... END" is the -# idiomatic idempotency-guard pattern (T-SQL has no native "CREATE TABLE IF -# NOT EXISTS"/"CREATE SCHEMA IF NOT EXISTS"), so without this exclusion -# nearly every T-SQL file would pick up a spurious, never-resolvable -# dependency on "sys.*" and get shuffled into the delayed-retry path, whose -# reverse-order resolution can then apply files out of their intended order. -_SYSTEM_SCHEMAS = {"pg_catalog", "information_schema", "sys"} +# SYSTEM_SCHEMAS (see dialect.py) matters most here for T-SQL, where +# "IF NOT EXISTS (SELECT ... FROM sys.schemas/sys.tables/sys.objects ...) +# BEGIN CREATE ... END" is the idiomatic idempotency-guard pattern (T-SQL has +# no native "CREATE TABLE IF NOT EXISTS"/"CREATE SCHEMA IF NOT EXISTS"), so +# without this exclusion nearly every T-SQL file would pick up a spurious, +# never-resolvable dependency on "sys.*" and get shuffled into the +# delayed-retry path, whose reverse-order resolution can then apply files +# out of their intended order. _DECLARE_RE = re.compile( r"CREATE\s+(?:OR\s+REPLACE\s+)?(?:TABLE|VIEW|FUNCTION|PROCEDURE|TYPE|SCHEMA)\s+(\w+\.\w+)", re.IGNORECASE @@ -91,7 +89,7 @@ def _get_sql_deps_regex_fallback(sql: str) -> set[str]: is fine.""" declares = set(_DECLARE_RE.findall(sql)) deps = set(_DEPEND_RE.findall(sql)) - deps = {d for d in deps if d.split(".", 1)[0].lower() not in _SYSTEM_SCHEMAS} + deps = {d for d in deps if d.split(".", 1)[0].lower() not in SYSTEM_SCHEMAS} return deps - declares @@ -111,7 +109,7 @@ def _get_sql_deps(sql: str, dialect: Dialect = POSTGRES) -> set[str]: for t in e.find_all(exp.Table): db_node = t.args.get("db") if t.args.get("this") is not None and db_node is not None: - if db_node.name.lower() in _SYSTEM_SCHEMAS: + if db_node.name.lower() in SYSTEM_SCHEMAS: continue # exp.table_name(), not str(t): str() includes " AS alias" for # an aliased reference (e.g. "FROM editing.visit v"), which diff --git a/tests/test_schemas.py b/tests/test_schemas.py index d80476f..5e385c0 100644 --- a/tests/test_schemas.py +++ b/tests/test_schemas.py @@ -54,6 +54,14 @@ def test_mssql_sys_schema_excluded(self): sql = "IF NOT EXISTS (SELECT 1 FROM sys.schemas WHERE name = 'app') BEGIN EXEC('CREATE SCHEMA app'); END" assert sql_schemas(sql, MSSQL) == frozenset({"app"}) + def test_create_schema_alongside_another_real_table_reference(self): + # Regression test: a bare `CREATE SCHEMA x` parses to a Table node + # with an empty name (its schema name lives in `db`, not `this`) -- + # the guard against sqlglot's EXEC(...) literal misparse (which has + # a *non-empty*, sentence-shaped name) must not also reject this. + sql = "CREATE SCHEMA analytics;\nSELECT 1 FROM public.tenants;\n" + assert sql_schemas(sql) == frozenset({"analytics", "public"}) + def test_unparseable_content_falls_back_to_regex(self): # Deliberately malformed SQL that still contains a schema-qualified # reference sqlglot can't make sense of as a whole statement. From f2d6c8d17443e9f6d58bf1c7bd2848a022122f90 Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Tue, 8 Sep 2026 06:33:49 +0200 Subject: [PATCH 3/5] Bump version to 0.5.0 for schema filtering New public surface (pgdevkit.schemas module, --schema/--exclude-schema on compare/migrate check/apply, --area/--exclude-area and --schema/--exclude-schema on testdb up/reset) warrants a minor bump, not a patch. auto-release.yml tags/releases/publishes automatically once this lands on main. Co-Authored-By: Claude Sonnet 5 Claude-Session: https://claude.ai/code/session_01AggszcqpDQAqPPEJCVVCLM --- pyproject.toml | 2 +- uv.lock | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 14d1801..cd49188 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,7 +11,7 @@ packages = ["pgdevkit"] [project] name = "pgdevkit" -version = "0.4.0" +version = "0.5.0" description = "A helper for developing with Postgres" readme = "README.md" requires-python = ">=3.14" diff --git a/uv.lock b/uv.lock index e9540a1..f6ad2a1 100644 --- a/uv.lock +++ b/uv.lock @@ -313,7 +313,7 @@ wheels = [ [[package]] name = "pgdevkit" -version = "0.4.0" +version = "0.5.0" source = { editable = "." } dependencies = [ { name = "docker" }, From de6c57dc9d1aac8036a3b6aca33e0858113500e4 Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Tue, 8 Sep 2026 06:37:35 +0200 Subject: [PATCH 4/5] Add regression tests locking in lazy area/schema detection Verified empirically that parse_directory() and _iter_sql_files() already skip parse_areas()/sql_schemas() entirely when that axis isn't being filtered on (each is gated by "(areas or exclude_areas) and ..." / "(schemas or exclude_schemas) and ..."), and that filter_by_area()/filter_by_schema() never touch the filesystem when neither only nor exclude is set -- no behavior change needed, this was already the case. Adds explicit call-counting tests so a future edit can't silently make either check unconditional. Co-Authored-By: Claude Sonnet 5 Claude-Session: https://claude.ai/code/session_01AggszcqpDQAqPPEJCVVCLM --- tests/test_parser_schemas.py | 54 +++++++++++++++++++++++++++ tests/testdb/test_schema_filtering.py | 53 ++++++++++++++++++++++++++ 2 files changed, 107 insertions(+) diff --git a/tests/test_parser_schemas.py b/tests/test_parser_schemas.py index fe96b56..13d6e4f 100644 --- a/tests/test_parser_schemas.py +++ b/tests/test_parser_schemas.py @@ -2,6 +2,7 @@ from pathlib import Path +import pgdevkit.parser as parser_mod from pgdevkit.parser import parse_directory @@ -42,3 +43,56 @@ def test_area_and_schema_filters_combine(self, tmp_path: Path): schema = parse_directory(tmp_path, areas=frozenset({"billing"}), schemas=frozenset({"reporting"})) assert set(schema.tables) == {"reporting.a"} + + +class TestParseDirectoryLaziness: + """Area/schema detection is real work (a regex scan; a full sqlglot parse, + for schema) done per file on top of the parse parse_directory needs + anyway -- so each check must run only when the axis it belongs to is + actually being filtered on, not unconditionally alongside the real parse.""" + + def test_no_filters_never_detects_areas_or_schemas(self, tmp_path: Path, monkeypatch): + _write(tmp_path, "a.sql", "CREATE TABLE billing.a (id int);\n") + calls: dict[str, int] = {"areas": 0, "schemas": 0} + monkeypatch.setattr( + parser_mod, "parse_areas", lambda content: calls.__setitem__("areas", calls["areas"] + 1) or frozenset() + ) + monkeypatch.setattr( + parser_mod, + "sql_schemas", + lambda content, dialect: calls.__setitem__("schemas", calls["schemas"] + 1) or frozenset(), + ) + + parse_directory(tmp_path) + + assert calls == {"areas": 0, "schemas": 0} + + def test_area_filter_alone_never_detects_schemas(self, tmp_path: Path, monkeypatch): + _write(tmp_path, "a.sql", "-- area: billing\nCREATE TABLE billing.a (id int);\n") + calls = 0 + + def counting_sql_schemas(content, dialect): + nonlocal calls + calls += 1 + return frozenset() + + monkeypatch.setattr(parser_mod, "sql_schemas", counting_sql_schemas) + + parse_directory(tmp_path, areas=frozenset({"billing"})) + + assert calls == 0 + + def test_schema_filter_alone_never_detects_areas(self, tmp_path: Path, monkeypatch): + _write(tmp_path, "a.sql", "CREATE TABLE billing.a (id int);\n") + calls = 0 + + def counting_parse_areas(content): + nonlocal calls + calls += 1 + return frozenset() + + monkeypatch.setattr(parser_mod, "parse_areas", counting_parse_areas) + + parse_directory(tmp_path, schemas=frozenset({"billing"})) + + assert calls == 0 diff --git a/tests/testdb/test_schema_filtering.py b/tests/testdb/test_schema_filtering.py index c8a4354..e842184 100644 --- a/tests/testdb/test_schema_filtering.py +++ b/tests/testdb/test_schema_filtering.py @@ -2,6 +2,7 @@ from pathlib import Path +import pgdevkit.testdb.schema as schema_mod from pgdevkit.testdb.schema import _iter_sql_files @@ -60,3 +61,55 @@ def test_area_and_schema_filters_combine(self, tmp_path: Path): result = _names(tmp_path, areas=frozenset({"billing"}), schemas=frozenset({"reporting"})) assert result == {Path("keep/tables/a.sql")} + + +class TestIterSqlFilesLaziness: + """Same guarantee as parse_directory's (see test_parser_schemas.py): area/ + schema detection must run only when that axis is actually being filtered + on, on top of the read+dependency-scan _iter_sql_files always does.""" + + def test_no_filters_never_detects_areas_or_schemas(self, tmp_path: Path, monkeypatch): + _write(tmp_path, "tables/x.sql", "CREATE TABLE billing.x (id int);\n") + calls: dict[str, int] = {"areas": 0, "schemas": 0} + monkeypatch.setattr( + schema_mod, "parse_areas", lambda content: calls.__setitem__("areas", calls["areas"] + 1) or frozenset() + ) + monkeypatch.setattr( + schema_mod, + "sql_schemas", + lambda content, dialect: calls.__setitem__("schemas", calls["schemas"] + 1) or frozenset(), + ) + + list(_iter_sql_files(tmp_path)) + + assert calls == {"areas": 0, "schemas": 0} + + def test_area_filter_alone_never_detects_schemas(self, tmp_path: Path, monkeypatch): + _write(tmp_path, "tables/x.sql", "-- area: billing\nCREATE TABLE billing.x (id int);\n") + calls = 0 + + def counting_sql_schemas(content, dialect): + nonlocal calls + calls += 1 + return frozenset() + + monkeypatch.setattr(schema_mod, "sql_schemas", counting_sql_schemas) + + list(_iter_sql_files(tmp_path, areas=frozenset({"billing"}))) + + assert calls == 0 + + def test_schema_filter_alone_never_detects_areas(self, tmp_path: Path, monkeypatch): + _write(tmp_path, "tables/x.sql", "CREATE TABLE billing.x (id int);\n") + calls = 0 + + def counting_parse_areas(content): + nonlocal calls + calls += 1 + return frozenset() + + monkeypatch.setattr(schema_mod, "parse_areas", counting_parse_areas) + + list(_iter_sql_files(tmp_path, schemas=frozenset({"billing"}))) + + assert calls == 0 From fc46ab7994da74e392164b2b65c875a0c731b565 Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Tue, 8 Sep 2026 06:44:40 +0200 Subject: [PATCH 5/5] Fix two schema-detection false results found by code review - sql_schemas()'s regex fallback scanned raw file text, so a schema merely *mentioned* in a "--" comment (e.g. "-- see billing.invoice for context") was wrongly detected as a real reference, defeating both --schema (falsely excluded) and --exclude-schema (falsely dropped) for such a file. Fixed by stripping comments first, via a comment-stripper extracted out of migrate.py into a new shared pgdevkit/sql_text.py module (schemas.py couldn't import migrate.py's copy directly -- migrate.py already imports schemas.py). - The guard against sqlglot's EXEC(...)-literal misparse required a table's name to fully match \w+, which also rejected legitimate quoted identifiers containing punctuation (e.g. billing."my-table"), silently dropping that table's schema. Loosened to reject only names containing whitespace, which still catches the original misparse (a full SQL statement embedded as literal text) without rejecting real quoted identifiers. Both reproduced and confirmed fixed against sql_schemas() directly; regression tests added for each, plus for the existing EXEC(...) and bare-CREATE-SCHEMA cases they abut. Co-Authored-By: Claude Sonnet 5 Claude-Session: https://claude.ai/code/session_01AggszcqpDQAqPPEJCVVCLM --- pgdevkit/migrate.py | 54 +++-------------------------------------- pgdevkit/schemas.py | 32 +++++++++++++++---------- pgdevkit/sql_text.py | 56 +++++++++++++++++++++++++++++++++++++++++++ tests/test_migrate.py | 4 ++-- tests/test_schemas.py | 18 ++++++++++++++ 5 files changed, 99 insertions(+), 65 deletions(-) create mode 100644 pgdevkit/sql_text.py diff --git a/pgdevkit/migrate.py b/pgdevkit/migrate.py index 6d5eae3..dfa186c 100644 --- a/pgdevkit/migrate.py +++ b/pgdevkit/migrate.py @@ -20,6 +20,7 @@ from .areas import filter_by_area from .schemas import filter_by_schema +from .sql_text import strip_line_comments _IDENTIFIER = r"[A-Za-z_][A-Za-z0-9_]*" _DEFAULT_TRACKING_TABLE = "public.schema_migrations" @@ -125,61 +126,12 @@ def _split_sql(sql: str) -> list[str]: return stmts -def _strip_line_comments(sql: str) -> str: - """Drop '--' line comments, respecting string literals and $$...$$ blocks.""" - buf: list[str] = [] - i = 0 - in_string = False - in_line_comment = False - dollar_tag: str | None = None - - while i < len(sql): - c = sql[i] - if in_line_comment: - if c == "\n": - in_line_comment = False - buf.append(c) - elif dollar_tag is not None: - buf.append(c) - if c == "$" and sql[i:i + len(dollar_tag)] == dollar_tag: - buf.extend(list(dollar_tag[1:])) - i += len(dollar_tag) - dollar_tag = None - continue - elif in_string: - if c == "'" and i + 1 < len(sql) and sql[i + 1] == "'": - buf.append(c) - buf.append(sql[i + 1]) - i += 2 - continue - elif c == "'": - in_string = False - buf.append(c) - elif c == "-" and i + 1 < len(sql) and sql[i + 1] == "-": - in_line_comment = True - elif c == "$": - m = re.match(r"\$([A-Za-z0-9_]*)\$", sql[i:]) - if m: - dollar_tag = m.group(0) - buf.extend(list(dollar_tag)) - i += len(dollar_tag) - continue - buf.append(c) - elif c == "'": - in_string = True - buf.append(c) - else: - buf.append(c) - i += 1 - return "".join(buf) - - def _created_table_names(stmts: list[str]) -> list[str]: """Names of tables any CREATE TABLE statement targets, parsed via sqlglot (falls back to regex on comment-stripped text for statements sqlglot's postgres dialect can't parse).""" names: list[str] = [] for stmt in stmts: - stripped = _strip_line_comments(stmt) + stripped = strip_line_comments(stmt) if not re.search(r"CREATE\s+TABLE", stripped, re.IGNORECASE): continue # skip sqlglot entirely for statements that can't be a CREATE TABLE try: @@ -221,7 +173,7 @@ def _idempotent_target(stmt: str) -> tuple[str, ...] | None: COLUMN. None if the statement isn't one of these shapes — including any CREATE OR REPLACE, which is never safe to treat as a no-op just because the object exists, since the migration could be replacing it with different content.""" - stripped = _strip_line_comments(stmt).strip() + stripped = strip_line_comments(stmt).strip() if re.search(r"\bOR\s+REPLACE\b", stripped, re.IGNORECASE): return None diff --git a/pgdevkit/schemas.py b/pgdevkit/schemas.py index 8224792..de831d6 100644 --- a/pgdevkit/schemas.py +++ b/pgdevkit/schemas.py @@ -24,6 +24,7 @@ import sqlglot.expressions as exp from .dialect import Dialect, POSTGRES, SYSTEM_SCHEMAS +from .sql_text import strip_line_comments _CREATE_SCHEMA_RE = re.compile( r"CREATE\s+SCHEMA\s+(?:IF\s+NOT\s+EXISTS\s+)?(?:AUTHORIZATION\s+)?\"?(\w+)\"?", re.IGNORECASE @@ -32,13 +33,18 @@ def _regex_fallback(sql: str) -> frozenset[str]: - """Crude schema scan used when sqlglot can't parse `sql` at all. Errs - toward over-matching (any `schema.name`-shaped token) rather than - under-matching -- a false positive here only keeps a file in scope for - one extra schema filter, where a false negative would silently drop it - from an `only` filter.""" - schemas = set(_CREATE_SCHEMA_RE.findall(sql)) - schemas.update(m.group(1) for m in _QUALIFIED_REF_RE.finditer(sql)) + """Crude schema scan used when sqlglot can't parse `sql` at all, or found + no reference. Errs toward over-matching (any `schema.name`-shaped token) + rather than under-matching -- a false positive here only keeps a file in + scope for one extra schema filter, where a false negative would silently + drop it from an `only` filter. Comments are stripped first so a schema + merely *mentioned* in one (e.g. "-- see billing.invoice for context") + isn't mistaken for a real reference; string literals are left as-is + (matching one inside a literal is the same over-matching tradeoff as + above, not worth the complexity of also blanking them out).""" + stripped = strip_line_comments(sql) + schemas = set(_CREATE_SCHEMA_RE.findall(stripped)) + schemas.update(m.group(1) for m in _QUALIFIED_REF_RE.finditer(stripped)) return frozenset(s for s in schemas if s.lower() not in SYSTEM_SCHEMAS) @@ -61,11 +67,13 @@ def sql_schemas(content: str, dialect: Dialect = POSTGRES) -> frozenset[str]: # whose "name" is the whole literal string, not a real # identifier -- e.g. a quoted "CREATE SCHEMA app". Guard against # treating that as a schema-qualified (or default-schema) table - # reference by requiring a *non-empty* name to actually look like - # one -- but still accept an empty name, which is the legitimate - # shape sqlglot gives a bare `CREATE SCHEMA x` (whose schema - # name ends up in `db`, not `this`). - if t.name and not re.fullmatch(r"\w+", t.name): + # reference by rejecting a name containing whitespace -- a real + # identifier, quoted or not (even one with hyphens or other + # punctuation), never has any, while embedded SQL text always + # does. Still accepts an empty name, the legitimate shape + # sqlglot gives a bare `CREATE SCHEMA x` (whose schema name ends + # up in `db`, not `this`). + if t.name and re.search(r"\s", t.name): continue db_node = t.args.get("db") name = db_node.name if db_node else dialect.default_schema diff --git a/pgdevkit/sql_text.py b/pgdevkit/sql_text.py new file mode 100644 index 0000000..4ba63a3 --- /dev/null +++ b/pgdevkit/sql_text.py @@ -0,0 +1,56 @@ +"""Small text-level SQL helpers shared by modules that can't import each +other directly (`migrate.py` imports `schemas.py`, so `schemas.py` can't +import back from `migrate.py`).""" + +from __future__ import annotations + +import re + + +def strip_line_comments(sql: str) -> str: + """Drop '--' line comments, respecting string literals and $$...$$ blocks.""" + buf: list[str] = [] + i = 0 + in_string = False + in_line_comment = False + dollar_tag: str | None = None + + while i < len(sql): + c = sql[i] + if in_line_comment: + if c == "\n": + in_line_comment = False + buf.append(c) + elif dollar_tag is not None: + buf.append(c) + if c == "$" and sql[i:i + len(dollar_tag)] == dollar_tag: + buf.extend(list(dollar_tag[1:])) + i += len(dollar_tag) + dollar_tag = None + continue + elif in_string: + if c == "'" and i + 1 < len(sql) and sql[i + 1] == "'": + buf.append(c) + buf.append(sql[i + 1]) + i += 2 + continue + elif c == "'": + in_string = False + buf.append(c) + elif c == "-" and i + 1 < len(sql) and sql[i + 1] == "-": + in_line_comment = True + elif c == "$": + m = re.match(r"\$([A-Za-z0-9_]*)\$", sql[i:]) + if m: + dollar_tag = m.group(0) + buf.extend(list(dollar_tag)) + i += len(dollar_tag) + continue + buf.append(c) + elif c == "'": + in_string = True + buf.append(c) + else: + buf.append(c) + i += 1 + return "".join(buf) diff --git a/tests/test_migrate.py b/tests/test_migrate.py index 778c394..832d395 100644 --- a/tests/test_migrate.py +++ b/tests/test_migrate.py @@ -4,9 +4,9 @@ _created_table_names, _idempotent_target, _split_sql, - _strip_line_comments, default_tracking_table, ) +from pgdevkit.sql_text import strip_line_comments def test_created_table_names_ignores_create_table_mentioned_in_a_comment(): @@ -52,7 +52,7 @@ def test_created_table_names_falls_back_to_regex_for_unparseable_statements(): def test_strip_line_comments_preserves_string_literals_containing_dashes(): sql = "SELECT '--not-a-comment' AS x -- a real comment\nFROM t;" - stripped = _strip_line_comments(sql) + stripped = strip_line_comments(sql) assert "--not-a-comment" in stripped assert "a real comment" not in stripped diff --git a/tests/test_schemas.py b/tests/test_schemas.py index 5e385c0..141d9bb 100644 --- a/tests/test_schemas.py +++ b/tests/test_schemas.py @@ -68,6 +68,24 @@ def test_unparseable_content_falls_back_to_regex(self): sql = "!!! not sql billing.invoice !!!" assert sql_schemas(sql) == frozenset({"billing"}) + def test_quoted_identifier_with_punctuation_still_detected(self): + # Regression test: the EXEC(...)-misparse guard used to require the + # whole table name to be `\w+`, which also rejected a legitimate + # quoted identifier containing a hyphen -- silently dropping the + # schema for any file using one. + sql = 'CREATE TABLE billing."my-table" (id int);' + assert sql_schemas(sql) == frozenset({"billing"}) + + def test_schema_mentioned_only_in_a_comment_is_not_detected(self): + # Regression test: sqlglot finds no real reference here, so this + # falls back to the regex scan -- which must not match "billing. + # invoice" inside the comment as if it were a real reference. The + # module's own contract (a file with no detectable reference is + # never excluded, and never excuses a required inclusion) depends on + # this file coming back as truly undetectable, not falsely "billing". + sql = "-- see billing.invoice for context\nSELECT 1;\n" + assert sql_schemas(sql) == frozenset() + class TestFileSchemas: def test_reads_from_disk(self, tmp_path: Path):