#!/usr/bin/env python3 """Copy afterhours-shifts from mgmt to prod. Dry-run unless --execute.""" from __future__ import annotations import argparse import sys import time import boto3 TABLE = "afterhours-shifts" SRC_ACCOUNT = "328440206208" DST_ACCOUNT = "011934824531" def _client(profile: str, region: str): session = boto3.Session(profile_name=profile, region_name=region) return session.client("dynamodb") def _scan_all(client, *, consistent: bool = False): items = [] kwargs = {"TableName": TABLE, "ConsistentRead": consistent} while True: resp = client.scan(**kwargs) items.extend(resp.get("Items", [])) start = resp.get("LastEvaluatedKey") if not start: return items kwargs["ExclusiveStartKey"] = start def _batch_write_all(client, items, *, sleep=time.sleep, max_attempts: int = 8) -> int: """Put every item. Retry UnprocessedItems with backoff. Raise if they remain.""" pending = [{"PutRequest": {"Item": item}} for item in items] written = 0 while pending: chunk, pending = pending[:25], pending[25:] to_send = chunk attempts = 0 while to_send: attempts += 1 if attempts > max_attempts: raise RuntimeError( f"batch_write_item left {len(to_send)} unprocessed after {max_attempts} attempts" ) resp = client.batch_write_item(RequestItems={TABLE: to_send}) unprocessed = resp.get("UnprocessedItems", {}).get(TABLE, []) written += len(to_send) - len(unprocessed) to_send = unprocessed if to_send: sleep(min(0.1 * (2 ** (attempts - 1)), 5.0)) return written def main() -> int: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--src-profile", required=True) parser.add_argument("--dst-profile", required=True) parser.add_argument("--region", default="us-east-1") parser.add_argument("--execute", action="store_true") args = parser.parse_args() src = _client(args.src_profile, args.region) dst = _client(args.dst_profile, args.region) src_id = boto3.Session(profile_name=args.src_profile).client("sts").get_caller_identity()["Account"] dst_id = boto3.Session(profile_name=args.dst_profile).client("sts").get_caller_identity()["Account"] if src_id != SRC_ACCOUNT: print(f"src account {src_id} is not mgmt {SRC_ACCOUNT}", file=sys.stderr) return 2 if dst_id != DST_ACCOUNT: print(f"dst account {dst_id} is not prod {DST_ACCOUNT}", file=sys.stderr) return 2 items = _scan_all(src, consistent=True) dst_count = dst.describe_table(TableName=TABLE)["Table"]["ItemCount"] print(f"src items={len(items)} dst describe ItemCount={dst_count}") if not args.execute: print("dry-run; pass --execute to BatchWriteItem") return 0 written = _batch_write_all(dst, items) after = _scan_all(dst, consistent=True) print(f"wrote={written} dst_scan={len(after)}") if len(after) != len(items): print("item counts differ after copy", file=sys.stderr) return 1 return 0 if __name__ == "__main__": raise SystemExit(main())