From 9e49ee79ad3ea36e1e8d0efebaf527d4434a7a4a Mon Sep 17 00:00:00 2001 From: hrodmn Date: Fri, 25 Sep 2026 06:46:50 -0500 Subject: [PATCH] chore: re-map collections in batches, tweak catalog backfill script --- README.md | 116 +++++- .../runtime/tests/test_backfill.py | 46 ++- .../runtime/tests/test_migrate_dps_cli.py | 312 ++++++++++++++++ .../tests/test_migrate_dps_collection_ids.py | 208 ++++++++++- scripts/backfill_dps_user_catalogs.py | 31 +- scripts/migrate_dps_collection_ids.py | 342 ++++++++++++------ 6 files changed, 916 insertions(+), 139 deletions(-) create mode 100644 cdk/constructs/DpsStacItemGenerator/runtime/tests/test_migrate_dps_cli.py diff --git a/README.md b/README.md index f524904..a1b395b 100644 --- a/README.md +++ b/README.md @@ -56,8 +56,9 @@ uv run --script scripts/backfill_dps_user_catalogs.py --apply ``` The backfill uses hydrated item metadata and actual collection IDs. It recognizes -both current three-part and legacy tag-specific four-part generated IDs, and -skips named, authorized, mixed, incomplete, and ambiguous collections. Historical +both current three-part and legacy tag-specific four-part generated IDs, whether +raw or slugified. Item tags may vary within a collection; username, algorithm +name, and version must agree. It skips named, authorized, mixed, incomplete, and ambiguous collections. Historical authorization cannot always be proven when its registry is incomplete, so review the dry-run report. Existing Collection metadata is preserved; apply only adds the parent relationship and creates a missing user Catalog. It does not rewrite @@ -68,7 +69,7 @@ then apply the database migration: ```bash ./scripts/migrate_dps_collection_ids.py --dry-run -./scripts/migrate_dps_collection_ids.py --apply +./scripts/migrate_dps_collection_ids.py --apply # default: --batch-size 100 ``` It recognizes four-part IDs (`username__algorithm__version__tag`), merges their @@ -78,6 +79,101 @@ legacy ID, but their Items still receive those metadata fields. For a deployed database, follow the [RDS connection guide](#connect-to-rds-through-an-ssm-tunnel) below and the RDS usage instructions in the migration script's docstring. +Conflict checks run one destination collection at a time rather than grouping +all migrating items together. Apply uses 100 source collections per transaction +by default; this default is not production-validated. Set `--batch-size` to a +positive integer after reviewing the dry-run plan. Destination groups are split +across batches as needed: the first chunk creates the destination collection +and later chunks append to it. The cap bounds source-collection count, not item +count or partition size. Retained collision sources that only receive metadata +are bounded by the same batching. + +Each transaction uses `work_mem=4MB`, `hash_mem_multiplier=1`, and disables +parallel query workers and JIT. These are per-operation budgets, not a total +memory limit; large groups can still need substantial temporary disk space and +time. Apply first materializes transformed items into temporary tables, then +inserts them into pgSTAC staging in a separate statement so partition- +maintenance triggers do not conflict with active reads. Completed batches remain +committed if a later batch fails; rerun `--apply` after fixing the failure to +finish remaining work. A failed current batch rolls back as one unit, including +source deletion and copying. + +Apply disables pgSTAC queueing only within each batch transaction: maintenance +runs before source partitions are deleted, avoiding queued references to deleted +partitions. Drain existing pgSTAC queued work before applying; the script does +not drain it or change the deployment-wide queue setting. Verify collection and +item counts and queued work after applying, then restore writers. After an RDS +out-of-memory restart, resolve any startup parameter errors and confirm the +instance is healthy before retrying, then monitor memory and storage. + +The migration is not safe with concurrent pgSTAC writes. Before `--apply`, +pause DPS generation, wait for its active invocations to finish, let the STAC +loader drain, and then pause the loader. The loader is the database writer; +disabling the DPS generator alone is not sufficient. Stop any other direct +pgSTAC writers as well. + +The generator's queued messages remain in SQS while its mapping is disabled. +Pause its SQS-to-Lambda mapping: + +```bash +STAGE=dev # change as appropriate +FUNCTION_NAME=$(aws cloudformation list-exports \ + --query "Exports[?Name=='dps-stac-item-generator-function-name-${STAGE}'].Value | [0]" \ + --output text) +GENERATOR_QUEUE_URL=$(aws cloudformation list-exports \ + --query "Exports[?Name=='dps-stac-item-generator-queue-url-${STAGE}'].Value | [0]" \ + --output text) +GENERATOR_MAPPING_UUID=$(aws lambda list-event-source-mappings \ + --function-name "$FUNCTION_NAME" \ + --query 'EventSourceMappings[0].UUID' \ + --output text) + +aws lambda update-event-source-mapping \ + --uuid "$GENERATOR_MAPPING_UUID" --no-enabled +aws lambda get-event-source-mapping --uuid "$GENERATOR_MAPPING_UUID" \ + --query '{State:State,StateTransitionReason:StateTransitionReason}' \ + --output table +``` + +Continue only after the mapping state is `Disabled` and the generator queue has +no in-flight messages: + +```bash +aws sqs get-queue-attributes --queue-url "$GENERATOR_QUEUE_URL" \ + --attribute-names ApproximateNumberOfMessagesNotVisible \ + --query 'Attributes.ApproximateNumberOfMessagesNotVisible' --output text +``` + +Let the loader finish its visible and in-flight messages. Find its Lambda +function in the deployed stack, set `LOADER_FUNCTION_NAME` to the physical ID +of the `stac-item-loader` Lambda, then disable its mapping: + +```bash +STACK="MAAP-STAC-${STAGE}-userSTAC" +aws cloudformation list-stack-resources --stack-name "$STACK" \ + --query 'StackResourceSummaries[?ResourceType==`AWS::Lambda::Function`].[LogicalResourceId,PhysicalResourceId]' \ + --output table + +LOADER_FUNCTION_NAME='' +LOADER_MAPPING_UUID=$(aws lambda list-event-source-mappings \ + --function-name "$LOADER_FUNCTION_NAME" \ + --query 'EventSourceMappings[0].UUID' \ + --output text) + +aws lambda update-event-source-mapping --uuid "$LOADER_MAPPING_UUID" --no-enabled +aws lambda get-event-source-mapping --uuid "$LOADER_MAPPING_UUID" \ + --query '{State:State,StateTransitionReason:StateTransitionReason}' \ + --output table +``` + +Run the migration only after the loader mapping state is `Disabled`. After the +migration and its verification, restore the loader first, then DPS generation: + +```bash +aws lambda update-event-source-mapping --uuid "$LOADER_MAPPING_UUID" --enabled +aws lambda update-event-source-mapping --uuid "$GENERATOR_MAPPING_UUID" --enabled +``` + Collection-only STAC transactions can still be enabled with: - `USER_STAC_COLLECTION_TRANSACTIONS_AUTH_MODE=basic` @@ -193,7 +289,7 @@ secrets belonging to that CDK deployment: ```bash aws sts get-caller-identity -STAGE=test # change as appropriate +STAGE=dev # change as appropriate STACK="MAAP-STAC-${STAGE}-userSTAC" # userSTAC or pgSTAC aws cloudformation list-stack-resources \ @@ -202,7 +298,7 @@ aws cloudformation list-stack-resources \ --output table ``` -You can also find these under **CloudFormation → stack → Resources**. +You can also find these under **CloudFormation → stack → Resources**. Select the database secret whose ID contains `pgstacdbbootstrappersecret`, not the STAC HTTP basic-auth secret. CloudFormation gives you the secret's identifier; retrieve its value @@ -231,7 +327,7 @@ endpoint from the same secret and start the session. Variables set in the first terminal are not available in this terminal: ```bash -STAGE=test # use the same stage as above +STAGE=dev # use the same stage as above TYPE=internal # use public for the pgSTAC stack SECRET_ID='' RDS_HOST=$(aws secretsmanager get-secret-value \ @@ -243,11 +339,17 @@ INSTANCE_ID=$(aws ssm get-parameter \ aws ssm start-session \ --target "$INSTANCE_ID" \ --document-name AWS-StartPortForwardingSessionToRemoteHost \ - --parameters "{\"host\":[\"$RDS_HOST\"],\"portNumber\":[\"5432\"],\"localPortNumber\":[\"15432\"]}" + --parameters "{\"host\":[\"$RDS_HOST\"],\"portNumber\":[\"5432\"],\"localPortNumber\":[\"15432\"]}" \ + --cli-read-timeout 0 ``` Leave this terminal open while you use the database. The EC2 host needs network access to RDS on port 5432, as it does for normal PgBouncer traffic. +Session Manager applies the account's idle and maximum-session-duration +preferences. Database traffic normally avoids the idle timeout, but the maximum +duration or a network interruption can still close the tunnel. For the batched +migration, the in-progress transaction rolls back while earlier batches remain +committed; restart the tunnel and rerun `--apply`. #### Connect with a local client diff --git a/cdk/constructs/DpsStacItemGenerator/runtime/tests/test_backfill.py b/cdk/constructs/DpsStacItemGenerator/runtime/tests/test_backfill.py index 471b363..a0e0b8f 100644 --- a/cdk/constructs/DpsStacItemGenerator/runtime/tests/test_backfill.py +++ b/cdk/constructs/DpsStacItemGenerator/runtime/tests/test_backfill.py @@ -48,7 +48,7 @@ def test_backfill_plan_is_dry_run_safe_and_idempotent(): def test_backfill_plans_legacy_tag_specific_collection(): - """The planner links legacy IDs when their tag matches item metadata.""" + """The planner links legacy IDs without assessing item tags.""" metadata = { "username": "alice", "algorithm_name": "algo", @@ -68,6 +68,50 @@ def test_backfill_plans_legacy_tag_specific_collection(): assert plan[0]["catalog_id"] == backfill.user_catalog_id("alice") +def test_backfill_allows_unslugified_collection_ids(): + """Ownership metadata can match a raw, mixed-case generated ID.""" + metadata = { + "username": "Alice", + "algorithm_name": "My Algorithm", + "algorithm_version": "1.0", + } + collection_id = backfill.COLLECTION_ID_FORMAT.format(**metadata) + records = { + collection_id: {"type": "Collection", "id": collection_id, "parent_ids": []} + } + + plan, skipped = backfill.build_plan([row(collection_id, **metadata)], records, {}) + + assert skipped == [] + assert plan[0]["catalog_id"] == backfill.user_catalog_id("Alice") + + +def test_backfill_ignores_mixed_and_missing_tags(): + """Tags do not affect ownership of a current generated collection.""" + metadata = { + "username": "alice", + "algorithm_name": "algo", + "algorithm_version": "1.0", + } + collection_id = backfill.generated_collection_id(metadata) + records = { + collection_id: {"type": "Collection", "id": collection_id, "parent_ids": []} + } + + plan, skipped = backfill.build_plan( + [ + row(collection_id, **metadata, tag="nightly"), + row(collection_id, **metadata, tag="release"), + row(collection_id, **metadata), + ], + records, + {}, + ) + + assert skipped == [] + assert plan[0]["catalog_id"] == backfill.user_catalog_id("alice") + + def test_backfill_skips_ambiguous_and_authorized_collections(): """The planner reports all conservative exclusions instead of guessing.""" generated_metadata = { diff --git a/cdk/constructs/DpsStacItemGenerator/runtime/tests/test_migrate_dps_cli.py b/cdk/constructs/DpsStacItemGenerator/runtime/tests/test_migrate_dps_cli.py new file mode 100644 index 0000000..47519cb --- /dev/null +++ b/cdk/constructs/DpsStacItemGenerator/runtime/tests/test_migrate_dps_cli.py @@ -0,0 +1,312 @@ +"""End-to-end migration checks in disposable databases cloned from pgSTAC.""" + +from __future__ import annotations + +import os +import subprocess +import sys +from pathlib import Path +from uuid import uuid4 + +import pytest +from psycopg import connect +from psycopg.conninfo import make_conninfo +from psycopg.sql import SQL, Identifier +from psycopg.types.json import Jsonb + +SCRIPT = Path(__file__).parents[5] / "scripts" / "migrate_dps_collection_ids.py" + + +@pytest.fixture +def database_url(): + """Clone the test database and drop the clone even when assertions fail.""" + template_url = os.environ.get("PGSTAC_TEST_DATABASE_URL") + if not template_url: + pytest.skip( + "Set PGSTAC_TEST_DATABASE_URL to an empty disposable pgSTAC database" + ) + with connect(template_url) as connection: + template = connection.info.dbname + assert ( + connection.execute("SELECT count(*) FROM pgstac.collections").fetchone()[0] + == 0 + ) + database = f"migration_cli_{uuid4().hex}" + with connect(template_url, dbname="template1", autocommit=True) as admin: + admin.execute( + SQL("CREATE DATABASE {} TEMPLATE {}").format( + Identifier(database), Identifier(template) + ) + ) + try: + yield make_conninfo(template_url, dbname=database) + finally: + admin.execute( + SQL("DROP DATABASE {} WITH (FORCE)").format(Identifier(database)) + ) + + +def snapshot(database_url): + """Read committed catalog contents and pending work from a new connection.""" + with connect(database_url) as connection: + return ( + connection.execute( + "SELECT content FROM pgstac.collections ORDER BY id" + ).fetchall(), + connection.execute( + "SELECT pgstac.format_item(items) FROM pgstac.items " + "ORDER BY collection, id" + ).fetchall(), + connection.execute( + "SELECT query FROM pgstac.query_queue ORDER BY query" + ).fetchall(), + ) + + +@pytest.mark.parametrize("use_queue", [True, False]) +def test_cli_batch_commit_rollback_and_maintenance(database_url, use_queue): + """Commit earlier batches, roll back a failed batch, then recover on rerun.""" + database_url = make_conninfo( + database_url, + options="-c search_path=pgstac,public " + f"-c pgstac.use_queue={str(use_queue).lower()} " + "-c enable_partition_pruning=off -c plan_cache_mode=force_generic_plan", + ) + existing = "aa__algo__1" + moving = existing + "__day" + retained = existing + "__night" + retained_again = existing + "__dawn" + new_target = "zz__algo__1" + new_sources = [new_target + "__day", new_target + "__night"] + collections = [ + existing, + moving, + retained, + retained_again, + *new_sources, + "unrelated", + ] + with connect(database_url) as connection: + for collection_id in collections: + connection.execute( + "SELECT pgstac.create_collection(%s)", + ( + Jsonb( + { + "type": "Collection", + "stac_version": "1.0.0", + "id": collection_id, + "description": collection_id, + "license": "proprietary", + "links": [], + "extent": { + "spatial": {"bbox": [[-180, -90, 180, 90]]}, + "temporal": {"interval": [[None, None]]}, + }, + } + ), + ), + ) + for collection_id, item_id, year in ( + (existing, "duplicate", 2020), + (moving, "moving", 2021), + (retained, "duplicate", 2022), + (retained_again, "duplicate", 2022), + (new_sources[0], "new-day", 2023), + (new_sources[1], "new-night", 2024), + ("unrelated", "untouched", 2025), + ): + connection.execute( + "INSERT INTO pgstac.items_staging_upsert (content) VALUES (%s)", + ( + Jsonb( + { + "type": "Feature", + "stac_version": "1.0.0", + "id": item_id, + "collection": collection_id, + "geometry": {"type": "Point", "coordinates": [0, 0]}, + "bbox": [0, 0, 0, 0], + "properties": { + "datetime": f"{year}-01-01T00:00:00Z", + "custom": 42, + }, + "links": [], + "assets": {"data": {"href": "s3://test/data.tif"}}, + } + ), + ), + ) + connection.execute( + "SELECT pgstac.update_partition_stats(partition) " + "FROM pgstac.partition_sys_meta WHERE collection = %s", + (existing,), + ) + # Record transaction-local settings on every committed source deletion. + connection.execute(""" + CREATE TABLE public.migration_settings_log ( + work_mem text, + hash_mem_multiplier text, + max_parallel_workers_per_gather text, + jit text, + use_queue text + ); + CREATE FUNCTION public.record_migration_settings() RETURNS trigger + LANGUAGE plpgsql AS $$ + BEGIN + INSERT INTO public.migration_settings_log + VALUES ( + current_setting('work_mem'), + current_setting('hash_mem_multiplier'), + current_setting('max_parallel_workers_per_gather'), + current_setting('jit'), + current_setting('pgstac.use_queue') + ); + RETURN OLD; + END $$; + CREATE TRIGGER record_migration_settings + AFTER DELETE ON pgstac.collections + FOR EACH ROW EXECUTE FUNCTION public.record_migration_settings(); + """) + # Fail in the second source-chunk batch, after the first commits. + connection.execute(""" + CREATE FUNCTION public.reject_late_delete() RETURNS trigger + LANGUAGE plpgsql AS $$ + BEGIN + RAISE EXCEPTION 'deliberate late migration failure'; + END $$; + CREATE TRIGGER reject_late_delete BEFORE DELETE ON pgstac.collections + FOR EACH ROW WHEN (OLD.id = 'zz__algo__1__day') + EXECUTE FUNCTION public.reject_late_delete(); + """) + # Start with no pending ingestion work, as required for production apply. + with connect(database_url, autocommit=True) as connection: + connection.execute("CALL pgstac.run_queued_queries()") + assert ( + connection.execute("SELECT count(*) FROM pgstac.query_queue").fetchone()[0] + == 0 + ) + assert ( + connection.execute( + "SELECT count(*) FROM pgstac.query_queue_history " + "WHERE error IS NOT NULL" + ).fetchone()[0] + == 0 + ) + before = snapshot(database_url) + command = [sys.executable, str(SCRIPT), "--database-url", database_url] + dry_run = subprocess.run( + [*command, "--dry-run", "--batch-size", "1"], + capture_output=True, + text=True, + timeout=120, + ) + assert dry_run.returncode == 0, dry_run.stderr + assert snapshot(database_url) == before + + failed = subprocess.run( + [*command, "--apply", "--batch-size", "1"], + capture_output=True, + text=True, + timeout=120, + ) + assert failed.returncode != 0 + assert "deliberate late migration failure" in failed.stderr + assert "Batch 2/5 failed and was rolled back" in failed.stderr + assert "Committed batch 1/5." in failed.stderr + assert "Committed batch 2/5." not in failed.stderr + partial = snapshot(database_url) + assert {row[0]["id"] for row in partial[0]} == { + existing, + retained, + retained_again, + new_target + "__day", + new_target + "__night", + "unrelated", + } + assert any( + item["collection"] == existing and item["id"] == "moving" + for (item,) in partial[1] + ) + assert all( + "maap-dps:tag" not in item["properties"] + for (item,) in partial[1] + if item["collection"] in (retained, retained_again) + ) + assert partial[2] == [] + with connect(database_url) as connection: + connection.execute("DROP TRIGGER reject_late_delete ON pgstac.collections") + connection.execute("DROP FUNCTION public.reject_late_delete()") + + applied = subprocess.run( + [*command, "--apply", "--batch-size", "1"], + capture_output=True, + text=True, + timeout=120, + ) + assert applied.returncode == 0, applied.stderr + after = snapshot(database_url) + assert after[2] == [], "Migration must not leave deferred partition maintenance" + with connect(database_url) as connection: + assert ( + connection.execute( + "SELECT pgstac.get_setting_bool('use_queue')" + ).fetchone()[0] + is use_queue + ) + settings = connection.execute( + "SELECT work_mem, hash_mem_multiplier, " + "max_parallel_workers_per_gather, jit, use_queue " + "FROM public.migration_settings_log" + ).fetchall() + assert settings + assert set(settings) == {("4MB", "1", "0", "off", "false")} + assert {row[0]["id"] for row in after[0]} == { + existing, + retained, + retained_again, + new_target, + "unrelated", + } + expected_items = [] + for (item,) in before[1]: + source = item["collection"] + if source in (moving, retained, retained_again, *new_sources): + username, algorithm, version, tag = source.split("__") + item["properties"].update( + { + "maap-dps:username": username, + "maap-dps:algorithm_name": algorithm, + "processing:version": version, + "maap-dps:tag": tag, + } + ) + if source not in (retained, retained_again): + item["collection"] = "__".join(source.split("__")[:-1]) + expected_items.append(item) + assert sorted( + (row[0] for row in after[1]), key=lambda item: (item["collection"], item["id"]) + ) == sorted(expected_items, key=lambda item: (item["collection"], item["id"])) + + # The procedure commits internally, so it must run outside a transaction. + with connect(database_url, autocommit=True) as connection: + connection.execute("CALL pgstac.run_queued_queries()") + assert ( + connection.execute("SELECT count(*) FROM pgstac.query_queue").fetchone()[0] + == 0 + ) + errors = connection.execute( + "SELECT query, error FROM pgstac.query_queue_history " + "WHERE error IS NOT NULL" + ).fetchall() + maintained = snapshot(database_url) + assert maintained[1] == after[1] + rerun = subprocess.run( + [*command, "--apply", "--batch-size", "1"], + capture_output=True, + text=True, + timeout=120, + ) + assert rerun.returncode == 0, rerun.stderr + assert snapshot(database_url)[:2] == maintained[:2] + assert not errors, errors diff --git a/cdk/constructs/DpsStacItemGenerator/runtime/tests/test_migrate_dps_collection_ids.py b/cdk/constructs/DpsStacItemGenerator/runtime/tests/test_migrate_dps_collection_ids.py index fa9359b..fb34b44 100644 --- a/cdk/constructs/DpsStacItemGenerator/runtime/tests/test_migrate_dps_collection_ids.py +++ b/cdk/constructs/DpsStacItemGenerator/runtime/tests/test_migrate_dps_collection_ids.py @@ -1,8 +1,14 @@ from __future__ import annotations import importlib.util +import os from pathlib import Path +import pytest +from psycopg import connect +from psycopg.rows import dict_row +from psycopg.types.json import Jsonb + SCRIPT = Path(__file__).parents[5] / "scripts" / "migrate_dps_collection_ids.py" MODULE_SPEC = importlib.util.spec_from_file_location( "migrate_dps_collection_ids", SCRIPT @@ -24,7 +30,7 @@ def __enter__(self) -> RecordingCursor: def __exit__(self, *args: object) -> None: return None - def execute(self, statement: str, parameters: tuple[list[str], ...]) -> None: + def execute(self, statement: str, parameters: tuple[list[str], ...] = ()) -> None: """Record an executed statement and its parameters.""" self.calls.append((statement, parameters)) @@ -40,6 +46,54 @@ def cursor(self) -> RecordingCursor: return self.recording_cursor +def test_migration_batches_share_destination_groups(): + """Destination groups share a batch up to the source cap.""" + plan = { + "x": ["x__one", "x__two"], + "y": ["y__one", "y__two"], + } + + batches = migration.migration_batches(plan, [], batch_size=4) + + assert batches == [(plan, [])] + + +def test_migration_batches_split_destination_groups(): + """A destination group larger than the cap is migrated in source chunks.""" + plan = {"x": ["x__one", "x__two"], "y": ["y__one", "y__two", "y__three"]} + + batches = migration.migration_batches( + plan, ["retained__algo__1__tag"], batch_size=2 + ) + + assert batches == [ + ({"x": ["x__one", "x__two"]}, []), + ({"y": ["y__one", "y__two"]}, []), + ({"y": ["y__three"]}, ["retained__algo__1__tag"]), + ] + + +def test_migration_batches_bound_metadata_only_sources(): + """Retained collision sources count toward the batch cap.""" + plan = {"x": ["x__one"]} + + batches = migration.migration_batches( + plan, ["retained__algo__1__a", "retained__algo__1__b"], batch_size=2 + ) + + assert batches == [ + ({"x": ["x__one"]}, ["retained__algo__1__a"]), + ({}, ["retained__algo__1__b"]), + ] + + +@pytest.mark.parametrize("value", ["0", "-1", "one"]) +def test_positive_batch_size_rejects_invalid_values(value): + """Batch size validation fails before a database connection is needed.""" + with pytest.raises(migration.argparse.ArgumentTypeError): + migration.positive_batch_size(value) + + def test_apply_item_metadata_keeps_legacy_collection_id(): """Conflicting legacy collections receive metadata without being moved.""" connection = RecordingConnection() @@ -48,6 +102,156 @@ def test_apply_item_metadata_keeps_legacy_collection_id(): migration.apply_item_metadata(connection, [source_id]) statement, parameters = connection.recording_cursor.calls[0] - assert "INSERT INTO pgstac.items_staging_upsert" in statement + assert "CREATE TEMP TABLE dps_metadata_items" in statement assert "'{collection}'" not in statement assert parameters == ([source_id], ["alice"], ["algorithm"], ["1.0"], ["nightly"]) + + +@pytest.mark.parametrize("use_queue", [True, False]) +def test_apply_with_pgstac_partition_triggers(use_queue): + """Preserve items while pgSTAC expands an existing destination partition.""" + database_url = os.environ.get("PGSTAC_TEST_DATABASE_URL") + if not database_url: + pytest.skip("Set PGSTAC_TEST_DATABASE_URL to a disposable pgSTAC database") + + target = "migration_test__algo__1" + source = f"{target}__day" + retained = f"{target}__night" + with ( + connect(database_url, row_factory=dict_row) as connection, + connection.transaction(force_rollback=True), + ): + connection.execute("SET LOCAL search_path = pgstac, public") + connection.execute( + "SELECT set_config('pgstac.use_queue', %s, true)", + (str(use_queue).lower(),), + ) + # Exercise the generic plan that can scan destination partitions too. + connection.execute("SET LOCAL plan_cache_mode = force_generic_plan") + connection.execute("SET LOCAL enable_partition_pruning = off") + connection.prepare_threshold = 0 + for collection_id in (target, source, retained): + connection.execute( + "SELECT pgstac.create_collection(%s)", + ( + Jsonb( + { + "type": "Collection", + "stac_version": "1.0.0", + "id": collection_id, + "description": "Migration regression test", + "license": "proprietary", + "links": [], + "extent": { + "spatial": {"bbox": [[-180, -90, 180, 90]]}, + "temporal": {"interval": [[None, None]]}, + }, + } + ), + ), + ) + for collection_id, item_id, date in ( + (target, "existing", "2020-01-01T00:00:00Z"), + (source, "moving", "2021-01-01T00:00:00Z"), + (retained, "existing", "2022-01-01T00:00:00Z"), + ): + connection.execute( + "INSERT INTO pgstac.items_staging_upsert (content) VALUES (%s)", + ( + Jsonb( + { + "type": "Feature", + "stac_version": "1.0.0", + "id": item_id, + "collection": collection_id, + "geometry": {"type": "Point", "coordinates": [0, 0]}, + "bbox": [0, 0, 0, 0], + "properties": {"datetime": date}, + "links": [], + "assets": {}, + } + ), + ), + ) + + connection.execute( + "SELECT pgstac.update_partition_stats(partition) " + "FROM pgstac.partition_sys_meta WHERE collection = %s", + (target,), + ) + migration.apply_item_metadata(connection, [retained]) + migration.apply_migration(connection, {target: [source]}) + + rows = connection.execute( + "SELECT collection, id, pgstac.format_item(items) AS content " + "FROM pgstac.items WHERE collection = ANY(%s) ORDER BY collection, id", + ([target, source, retained],), + ).fetchall() + assert [(row["collection"], row["id"]) for row in rows] == [ + (target, "existing"), + (target, "moving"), + (retained, "existing"), + ] + assert rows[1]["content"]["properties"]["maap-dps:tag"] == "day" + assert rows[2]["content"]["properties"]["maap-dps:tag"] == "night" + assert rows[1]["content"]["properties"]["datetime"] == "2021-01-01T00:00:00Z" + assert ( + connection.execute( + "SELECT id FROM pgstac.collections WHERE id = %s", (source,) + ).fetchone() + is None + ) + + +def test_conflicts_across_collection_groups(): + """Execute conflict checks against an empty disposable PostgreSQL database.""" + database_url = os.environ.get("MIGRATION_TEST_DATABASE_URL") + if not database_url: + pytest.skip("Set MIGRATION_TEST_DATABASE_URL to an empty test database") + + with ( + connect(database_url, row_factory=dict_row) as connection, + connection.transaction(force_rollback=True), + ): + # Deliberately fail if pgstac exists; never replace real catalog tables. + connection.execute("CREATE SCHEMA pgstac") + connection.execute("CREATE TABLE pgstac.collections (id text PRIMARY KEY)") + connection.execute( + "CREATE TABLE pgstac.items (collection text, id text) " + "PARTITION BY LIST (collection)" + ) + connection.execute( + "CREATE TABLE pgstac.items_default PARTITION OF pgstac.items DEFAULT" + ) + connection.execute( + """ + INSERT INTO pgstac.collections VALUES + ('a__algo__1'), ('a__algo__1__day'), ('a__algo__1__night'), + ('a__algo__1__safe'), ('b__algo__1__day'), ('b__algo__1__night'), + ('c__algo__1__empty'), ('unrelated'); + INSERT INTO pgstac.items VALUES + ('a__algo__1__day', 'sibling-collision'), + ('a__algo__1__night', 'sibling-collision'), + ('a__algo__1', 'target-collision'), + ('a__algo__1__day', 'target-collision'), + ('a__algo__1__safe', 'safe'), + ('b__algo__1__day', 'sibling-collision'), + ('b__algo__1__night', 'safe'), + ('unrelated', 'safe'); + """ + ) + plan = migration.migration_plan(connection) + assert migration.conflicting_source_collections(connection, {}) == [] + assert migration.conflicting_item_ids(connection, {}) == [] + skipped = migration.conflicting_source_collections(connection, plan) + assert skipped == ["a__algo__1__day", "a__algo__1__night"] + assert migration.conflicting_item_ids(connection, plan) == [ + "a__algo__1/sibling-collision", + "a__algo__1/target-collision", + ] + remaining = { + target: [source for source in sources if source not in skipped] + for target, sources in plan.items() + } + assert migration.conflicting_source_collections(connection, remaining) == [] + assert migration.conflicting_item_ids(connection, remaining) == [] diff --git a/scripts/backfill_dps_user_catalogs.py b/scripts/backfill_dps_user_catalogs.py index 380fa7b..384bbd8 100755 --- a/scripts/backfill_dps_user_catalogs.py +++ b/scripts/backfill_dps_user_catalogs.py @@ -11,10 +11,12 @@ The default is a dry run. Review the report, then rerun with ``--apply``. The backfill uses hydrated item metadata and the actual collection ID; it does not parse collection IDs to infer ownership. It only handles collections whose -items agree on one complete DPS metadata tuple and whose ID matches either the -current generator default or its legacy tag-specific format. Collections -authorized by the supplied registry, named collections, mixed collections, and -incomplete or ambiguous metadata are reported and skipped. +items agree on the DPS ownership metadata (username, algorithm name, and +version) and whose ID matches either the current generator default or its +legacy tag-specific format, in raw or slugified form. Item tags are ignored: +they may vary within a collection. Collections authorized by the supplied +registry, named collections, mixed collections, and incomplete or ambiguous +metadata are reported and skipped. Historical authorization is not present in every item record. A collection that happens to have a generated-looking ID can therefore be indistinguishable @@ -49,11 +51,10 @@ DEFAULT_DATABASE_URL = "postgresql://username:password@127.0.0.1:5439/postgis" COLLECTION_ID_FORMAT = "{username}__{algorithm_name}__{algorithm_version}" LEGACY_COLLECTION_ID_FORMAT = "{username}__{algorithm_name}__{algorithm_version}__{tag}" -METADATA_FIELDS = ( +OWNERSHIP_METADATA_FIELDS = ( "username", "algorithm_name", "algorithm_version", - "tag", ) @@ -118,8 +119,7 @@ def collection_rows(connection: Any) -> list[dict[str, Any]]: pgstac.format_item(items)->'properties'->>'maap-dps:algorithm_name' AS algorithm_name, pgstac.format_item(items)->'properties'->>'processing:version' - AS algorithm_version, - pgstac.format_item(items)->'properties'->>'maap-dps:tag' AS tag + AS algorithm_version FROM pgstac.collections AS collections JOIN pgstac.items AS items ON items.collection = collections.id @@ -152,22 +152,25 @@ def build_plan( for collection_id, item_rows in grouped.items(): values = { field: {row.get(field) for row in item_rows if row.get(field)} - for field in METADATA_FIELDS + for field in OWNERSHIP_METADATA_FIELDS } - if any(not values[field] for field in METADATA_FIELDS): + if any(not values[field] for field in OWNERSHIP_METADATA_FIELDS): skipped.append((collection_id, "missing DPS metadata")) continue - if any(len(values[field]) != 1 for field in METADATA_FIELDS): + if any(len(values[field]) != 1 for field in OWNERSHIP_METADATA_FIELDS): skipped.append((collection_id, "mixed DPS metadata")) continue - metadata = {field: values[field].pop() for field in METADATA_FIELDS} + metadata = {field: values[field].pop() for field in OWNERSHIP_METADATA_FIELDS} username = metadata["username"] generated_ids = { + COLLECTION_ID_FORMAT.format(**metadata), generated_collection_id(metadata), - generated_collection_id(metadata, LEGACY_COLLECTION_ID_FORMAT), } - if collection_id not in generated_ids: + legacy_id, separator, legacy_tag = collection_id.rpartition("__") + if collection_id not in generated_ids and ( + not separator or not legacy_tag or legacy_id not in generated_ids + ): skipped.append((collection_id, "named or non-generated collection ID")) continue if is_authorized(username, collection_id, registry): diff --git a/scripts/migrate_dps_collection_ids.py b/scripts/migrate_dps_collection_ids.py index 43304a1..ee1ae77 100755 --- a/scripts/migrate_dps_collection_ids.py +++ b/scripts/migrate_dps_collection_ids.py @@ -21,13 +21,23 @@ uv run --script scripts/migrate_dps_collection_ids.py --database-url "" --dry-run uv run --script scripts/migrate_dps_collection_ids.py --database-url "" --apply + uv run --script scripts/migrate_dps_collection_ids.py --database-url "" \\ + --apply --batch-size 100 The empty --database-url tells psycopg to use the PG* environment variables. -Omitting it uses DATABASE_URL or the local Compose default instead. +Omitting it uses DATABASE_URL or the local Compose default instead. Apply uses +100 source collections per transaction by default; this default is not +production-validated. Destination groups are split across batches when needed; +the first chunk creates the destination collection and later chunks append to +it. Completed batches remain committed if a later batch fails; rerun after +fixing the failure to complete the remaining work. Metadata-only updates for +retained collision sources are batched too. Before applying, review the dry-run plan, confirm a recoverable backup, pause -writers, and drain in-flight ingestion. The conflict check does not prevent -concurrent writes. Verify collection and item counts and check pgSTAC queued -work before resuming ingestion; the deployed stack enables use_queue. +all writers, and drain in-flight ingestion and existing pgSTAC queued work. +Apply disables queueing only within each batch transaction. The conflict check +does not prevent concurrent writes. Verify collection and item counts and check +pgSTAC queued work before resuming ingestion; the deployed stack enables +use_queue. """ from __future__ import annotations @@ -43,6 +53,7 @@ LOGGER = logging.getLogger(__name__) DEFAULT_DATABASE_URL = "postgresql://username:password@127.0.0.1:5439/postgis" +DEFAULT_BATCH_SIZE = 100 def target_collection_id(collection_id: str) -> str | None: @@ -70,80 +81,99 @@ def conflicting_source_collections( connection: Any, plan: dict[str, list[str]] ) -> list[str]: """Return legacy collections containing items that conflict after merging.""" - sources = [source for source_ids in plan.values() for source in source_ids] - if not sources: - return [] - + skipped: set[str] = set() with connection.cursor() as cursor: - cursor.execute( - """ - WITH migrated_items AS ( - SELECT - COALESCE(mapping.target_id, items.collection) AS target_id, - items.collection AS source_id, - items.id + for target_id, source_ids in plan.items(): + LOGGER.info("Checking item-ID conflicts for %s", target_id) + cursor.execute( + """ + WITH conflicts AS ( + SELECT id + FROM pgstac.items + WHERE collection = ANY(%s) + GROUP BY id + HAVING count(*) > 1 + ) + SELECT DISTINCT items.collection AS source_id FROM pgstac.items - LEFT JOIN unnest(%s::text[], %s::text[]) - AS mapping(source_id, target_id) - ON items.collection = mapping.source_id + JOIN conflicts USING (id) WHERE items.collection = ANY(%s) - OR items.collection = ANY(%s) - ), conflicts AS ( - SELECT target_id, id - FROM migrated_items - GROUP BY target_id, id - HAVING count(*) > 1 + """, + ([target_id, *source_ids], source_ids), ) - SELECT DISTINCT migrated_items.source_id - FROM migrated_items - JOIN conflicts USING (target_id, id) - WHERE migrated_items.source_id = ANY(%s) - ORDER BY migrated_items.source_id - """, - ( - sources, - [target for target, source_ids in plan.items() for _ in source_ids], - sources, - list(plan), - sources, - ), - ) - return [row["source_id"] for row in cursor.fetchall()] + skipped.update(row["source_id"] for row in cursor.fetchall()) + return sorted(skipped) def conflicting_item_ids(connection: Any, plan: dict[str, list[str]]) -> list[str]: """Return item-ID conflicts that would be created by the migration.""" - sources = [source for source_ids in plan.values() for source in source_ids] - if not sources: - return [] - + conflicts: list[str] = [] with connection.cursor() as cursor: - cursor.execute( - """ - SELECT target_id, id - FROM ( - SELECT - COALESCE(mapping.target_id, items.collection) AS target_id, - items.id + for target_id, source_ids in sorted(plan.items()): + cursor.execute( + """ + SELECT id FROM pgstac.items - LEFT JOIN unnest(%s::text[], %s::text[]) - AS mapping(source_id, target_id) - ON items.collection = mapping.source_id - WHERE items.collection = ANY(%s) - OR items.collection = ANY(%s) - ) AS migrated_items - GROUP BY target_id, id - HAVING count(*) > 1 - ORDER BY target_id, id - """, - ( - sources, - [target for target, source_ids in plan.items() for _ in source_ids], - sources, - list(plan), - ), - ) - return [f"{row['target_id']}/{row['id']}" for row in cursor.fetchall()] + WHERE collection = ANY(%s) + GROUP BY id + HAVING count(*) > 1 + ORDER BY id + """, + ([target_id, *source_ids],), + ) + conflicts.extend(f"{target_id}/{row['id']}" for row in cursor.fetchall()) + return conflicts + + +def migration_batches( + plan: dict[str, list[str]], skipped_sources: list[str], batch_size: int +) -> list[tuple[dict[str, list[str]], list[str]]]: + """Split migration and metadata work into source-collection-sized batches.""" + groups = [ + ({target_id: source_ids[start : start + batch_size]}, []) + for target_id, source_ids in sorted(plan.items()) + for start in range(0, len(source_ids), batch_size) + ] + groups.extend(({}, [source_id]) for source_id in skipped_sources) + + batches: list[tuple[dict[str, list[str]], list[str]]] = [] + batch_plan: dict[str, list[str]] = {} + batch_skipped: list[str] = [] + source_count = 0 + for group_plan, group_skipped in groups: + group_size = sum(map(len, group_plan.values())) + len(group_skipped) + if source_count and source_count + group_size > batch_size: + batches.append((batch_plan, batch_skipped)) + batch_plan, batch_skipped, source_count = {}, [], 0 + batch_plan.update(group_plan) + batch_skipped.extend(group_skipped) + source_count += group_size + if batch_plan or batch_skipped: + batches.append((batch_plan, batch_skipped)) + return batches + + +def configure_transaction(connection: Any, *, disable_queue: bool) -> None: + """Apply the migration's per-transaction resource and queue settings.""" + connection.execute("SET LOCAL work_mem = '4MB'") + connection.execute("SET LOCAL hash_mem_multiplier = 1") + connection.execute("SET LOCAL max_parallel_workers_per_gather = 0") + connection.execute("SET LOCAL jit = off") + if disable_queue: + connection.execute("SET LOCAL pgstac.use_queue = false") + + +def positive_batch_size(value: str) -> int: + """Parse a strictly positive batch size for the command line.""" + try: + batch_size = int(value) + except ValueError as error: + raise argparse.ArgumentTypeError( + "batch size must be a positive integer" + ) from error + if batch_size <= 0: + raise argparse.ArgumentTypeError("batch size must be a positive integer") + return batch_size def apply_item_metadata(connection: Any, source_ids: list[str]) -> None: @@ -155,7 +185,7 @@ def apply_item_metadata(connection: Any, source_ids: list[str]) -> None: with connection.cursor() as cursor: cursor.execute( """ - INSERT INTO pgstac.items_staging_upsert (content) + CREATE TEMP TABLE dps_metadata_items (content) ON COMMIT DROP AS SELECT jsonb_set( jsonb_set( jsonb_set( @@ -186,11 +216,20 @@ def apply_item_metadata(connection: Any, source_ids: list[str]) -> None: [parts[3] for parts in source_parts], ), ) + # Finish reading items before pgSTAC's triggers alter their partitions. + cursor.execute( + "INSERT INTO pgstac.items_staging_upsert (content) " + "SELECT content FROM pg_temp.dps_metadata_items" + ) + cursor.execute("DROP TABLE pg_temp.dps_metadata_items") def apply_migration(connection: Any, plan: dict[str, list[str]]) -> None: """Create tag-free collections, move their items, and remove old collections.""" with connection.cursor() as cursor: + cursor.execute( + "CREATE TEMP TABLE dps_migration_items (content jsonb) ON COMMIT DROP" + ) for target_id, source_ids in plan.items(): source_id = source_ids[0] cursor.execute( @@ -206,7 +245,7 @@ def apply_migration(connection: Any, plan: dict[str, list[str]]) -> None: source_parts = [source_id.split("__") for source_id in source_ids] cursor.execute( """ - INSERT INTO pgstac.items_staging_upsert (content) + INSERT INTO pg_temp.dps_migration_items (content) SELECT jsonb_set( jsonb_set( jsonb_set( @@ -253,12 +292,19 @@ def apply_migration(connection: Any, plan: dict[str, list[str]]) -> None: [parts[3] for parts in source_parts], ), ) + # A separate statement releases the read's active partition scans. + cursor.execute( + "INSERT INTO pgstac.items_staging_upsert (content) " + "SELECT content FROM pg_temp.dps_migration_items" + ) + cursor.execute("TRUNCATE pg_temp.dps_migration_items") cursor.execute( "DELETE FROM pgstac.items WHERE collection = ANY(%s)", (source_ids,) ) cursor.execute( "DELETE FROM pgstac.collections WHERE id = ANY(%s)", (source_ids,) ) + cursor.execute("DROP TABLE pg_temp.dps_migration_items") def parse_args() -> argparse.Namespace: @@ -275,66 +321,132 @@ def parse_args() -> argparse.Namespace: parser.add_argument( "--dry-run", action="store_true", help="Report changes without applying them." ) + parser.add_argument( + "--batch-size", + type=positive_batch_size, + default=DEFAULT_BATCH_SIZE, + help=( + "Maximum source collections per apply transaction; destination groups " + "are split across batches (default: 100)." + ), + ) return parser.parse_args() def main() -> None: """Report or apply the DPS collection-ID migration.""" args = parse_args() - logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s") - - with connect(args.database_url, row_factory=dict_row) as connection: - plan = migration_plan(connection) - if not plan: - LOGGER.info("No four-part legacy DPS collection IDs found.") - return - - skipped_sources = conflicting_source_collections(connection, plan) - if skipped_sources: - LOGGER.warning( - "Skipping %d legacy collection(s) with duplicate item IDs: %s", - len(skipped_sources), - ", ".join(skipped_sources), - ) - skipped = set(skipped_sources) - plan = { - target_id: [ - source_id for source_id in source_ids if source_id not in skipped - ] - for target_id, source_ids in plan.items() - } - plan = { - target_id: source_ids - for target_id, source_ids in plan.items() - if source_ids - } - - for source_id in skipped_sources: - LOGGER.info("%s -> retain collection and add DPS item metadata", source_id) - for target_id, source_ids in plan.items(): - LOGGER.info("%s -> %s", ", ".join(source_ids), target_id) + logging.basicConfig( + level=logging.INFO, + format="%(asctime)s %(levelname)s %(message)s", + datefmt="%Y-%m-%dT%H:%M:%S%z", + ) + + # Autocommit keeps planning and each apply batch in explicit, short + # transactions instead of leaving a planning transaction open. + with connect( + args.database_url, row_factory=dict_row, autocommit=True + ) as connection: + with connection.transaction(): + configure_transaction(connection, disable_queue=False) + plan = migration_plan(connection) + if not plan: + LOGGER.info("No four-part legacy DPS collection IDs found.") + return + + skipped_sources = conflicting_source_collections(connection, plan) + if skipped_sources: + LOGGER.warning( + "Skipping %d legacy collection(s) with duplicate item IDs: %s", + len(skipped_sources), + ", ".join(skipped_sources), + ) + skipped = set(skipped_sources) + plan = { + target_id: [ + source_id + for source_id in source_ids + if source_id not in skipped + ] + for target_id, source_ids in plan.items() + } + plan = { + target_id: source_ids + for target_id, source_ids in plan.items() + if source_ids + } + + for source_id in skipped_sources: + LOGGER.info( + "%s -> retain collection and add DPS item metadata", source_id + ) + for target_id, source_ids in plan.items(): + LOGGER.info("%s -> %s", ", ".join(source_ids), target_id) + + conflicts = conflicting_item_ids(connection, plan) + if conflicts: + raise SystemExit( + "Refusing to merge duplicate item IDs: " + ", ".join(conflicts) + ) - conflicts = conflicting_item_ids(connection, plan) - if conflicts: - raise SystemExit( - "Refusing to merge duplicate item IDs: " + ", ".join(conflicts) + batches = migration_batches(plan, skipped_sources, args.batch_size) + LOGGER.info( + "Proposed %d batch(es), up to %d source collection(s) each.", + len(batches), + args.batch_size, ) - if args.dry_run or not args.apply: + for batch_number, (batch_plan, batch_skipped) in enumerate(batches, 1): + source_count = sum(map(len, batch_plan.values())) + len(batch_skipped) + LOGGER.info( + "Batch %d/%d: %d source collection(s), %d destination group(s)", + batch_number, + len(batches), + source_count, + len(batch_plan), + ) + + if args.dry_run or not args.apply: + LOGGER.info( + "Dry run. Re-run with --apply to migrate %d collection(s) and add " + "DPS item metadata to %d retained collection(s).", + sum(map(len, plan.values())), + len(skipped_sources), + ) + return + + for batch_number, (batch_plan, batch_skipped) in enumerate(batches, 1): + source_count = sum(map(len, batch_plan.values())) + len(batch_skipped) LOGGER.info( - "Dry run. Re-run with --apply to migrate %d collection(s) and add " - "DPS item metadata to %d retained collection(s).", - sum(map(len, plan.values())), - len(skipped_sources), + "Applying batch %d/%d: %d source collection(s), " + "%d destination group(s)", + batch_number, + len(batches), + source_count, + len(batch_plan), ) - return + try: + with connection.transaction(): + configure_transaction(connection, disable_queue=True) + apply_item_metadata(connection, batch_skipped) + apply_migration(connection, batch_plan) + except Exception: + LOGGER.exception( + "Batch %d/%d failed and was rolled back; %d earlier batch(es) " + "remain committed. Restore the connection, fix the error, " + "and rerun.", + batch_number, + len(batches), + batch_number - 1, + ) + raise + LOGGER.info("Committed batch %d/%d.", batch_number, len(batches)) - apply_item_metadata(connection, skipped_sources) - apply_migration(connection, plan) LOGGER.info( "Migrated %d collection(s) and added DPS item metadata to %d retained " - "collection(s).", + "collection(s) across %d committed batch(es).", sum(map(len, plan.values())), len(skipped_sources), + len(batches), )