diff --git a/scripts/cutover/copy_dynamodb.py b/scripts/cutover/copy_dynamodb.py index d381f78..d476e2f 100644 --- a/scripts/cutover/copy_dynamodb.py +++ b/scripts/cutover/copy_dynamodb.py @@ -5,10 +5,10 @@ from __future__ import annotations import argparse import sys +import time import boto3 - TABLE = "afterhours-shifts" SRC_ACCOUNT = "328440206208" DST_ACCOUNT = "011934824531" @@ -19,9 +19,9 @@ def _client(profile: str, region: str): return session.client("dynamodb") -def _scan_all(client): +def _scan_all(client, *, consistent: bool = False): items = [] - kwargs = {"TableName": TABLE} + kwargs = {"TableName": TABLE, "ConsistentRead": consistent} while True: resp = client.scan(**kwargs) items.extend(resp.get("Items", [])) @@ -31,6 +31,29 @@ def _scan_all(client): 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) @@ -50,25 +73,15 @@ def main() -> int: print(f"dst account {dst_id} is not prod {DST_ACCOUNT}", file=sys.stderr) return 2 - items = _scan_all(src) + 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 = 0 - batch = [] - for item in items: - batch.append({"PutRequest": {"Item": item}}) - if len(batch) == 25: - dst.batch_write_item(RequestItems={TABLE: batch}) - written += len(batch) - batch = [] - if batch: - dst.batch_write_item(RequestItems={TABLE: batch}) - written += len(batch) - after = _scan_all(dst) + 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) diff --git a/scripts/cutover/recreate_holiday_schedules.py b/scripts/cutover/recreate_holiday_schedules.py index dcf5cbc..0308d1c 100644 --- a/scripts/cutover/recreate_holiday_schedules.py +++ b/scripts/cutover/recreate_holiday_schedules.py @@ -9,8 +9,10 @@ role. Dry-run unless --execute. from __future__ import annotations import argparse +import re import sys from datetime import datetime, timezone +from zoneinfo import ZoneInfo import boto3 from botocore.exceptions import ClientError @@ -20,6 +22,7 @@ DST_ACCOUNT = "011934824531" PROD_ROUTER_ARN = "arn:aws:lambda:us-east-1:011934824531:function:afterhours-holiday-router" PROD_ROLE_ARN = "arn:aws:iam::011934824531:role/tf-managed/afterhours-shift-manager-holiday-scheduler" PREFIXES = ("holiday-activate-", "holiday-deactivate-") +_AT = re.compile(r"^at\((\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2})\)$") def _client(profile: str, region: str): @@ -30,6 +33,22 @@ def _account(profile: str) -> str: return boto3.Session(profile_name=profile).client("sts").get_caller_identity()["Account"] +def schedule_when(detail: dict) -> datetime | None: + """UTC instant the one-off schedule fires, or None if it cannot be parsed.""" + expr = (detail.get("ScheduleExpression") or "").strip() + tzname = detail.get("ScheduleExpressionTimezone") or "America/New_York" + match = _AT.match(expr) + if match: + naive = datetime.strptime(match.group(1), "%Y-%m-%dT%H:%M:%S") + return naive.replace(tzinfo=ZoneInfo(tzname)).astimezone(timezone.utc) + at = detail.get("EndDate") or detail.get("StartDate") + if at is None: + return None + if at.tzinfo is None: + return at.replace(tzinfo=timezone.utc) + return at.astimezone(timezone.utc) + + def _list_holiday(client): names = [] token = None @@ -67,14 +86,15 @@ def main() -> int: now = datetime.now(timezone.utc) created = 0 skipped = 0 + failed = 0 for name in _list_holiday(src): detail = src.get_schedule(Name=name, GroupName="default") expr = detail.get("ScheduleExpression", "") tzname = detail.get("ScheduleExpressionTimezone", "America/New_York") - at = detail.get("EndDate") or detail.get("StartDate") - if at is not None and at < now: - print(f"skip past {name}") + when = schedule_when(detail) + if when is not None and when < now: + print(f"skip past {name} expr={expr}") skipped += 1 continue payload = { @@ -99,15 +119,20 @@ def main() -> int: dst.create_schedule(**payload) created += 1 except ClientError as exc: - if exc.response["Error"]["Code"] == "ConflictException": + code = exc.response["Error"]["Code"] + if code == "ConflictException": print(f"exists {name}") + elif code == "ValidationException": + print(f"skip invalid {name}: {exc.response['Error'].get('Message', code)}") + skipped += 1 else: - raise + print(f"failed {name}: {code}", file=sys.stderr) + failed += 1 - print(f"created={created} skipped_past={skipped} execute={args.execute}") + print(f"created={created} skipped_past={skipped} failed={failed} execute={args.execute}") if not args.execute: print("dry-run; pass --execute to CreateSchedule") - return 0 + return 1 if failed else 0 if __name__ == "__main__": diff --git a/tests/scripts/test_copy_dynamodb.py b/tests/scripts/test_copy_dynamodb.py new file mode 100644 index 0000000..90f4ca9 --- /dev/null +++ b/tests/scripts/test_copy_dynamodb.py @@ -0,0 +1,64 @@ +"""copy_dynamodb retries UnprocessedItems instead of counting them as written.""" + +import importlib.util +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[2] + + +def _load(): + spec = importlib.util.spec_from_file_location( + "copy_dynamodb", ROOT / "scripts" / "cutover" / "copy_dynamodb.py" + ) + mod = importlib.util.module_from_spec(spec) + sys.modules["copy_dynamodb"] = mod + spec.loader.exec_module(mod) + return mod + + +mod = _load() + + +class FakeDdb: + def __init__(self, unprocessed_first=None): + self.calls = [] + self.unprocessed_first = list(unprocessed_first or []) + self._first = True + + def batch_write_item(self, RequestItems): + batch = RequestItems[mod.TABLE] + self.calls.append(batch) + if self._first and self.unprocessed_first: + self._first = False + leftover = [req for req in batch if req in self.unprocessed_first] + return {"UnprocessedItems": {mod.TABLE: leftover} if leftover else {}} + return {"UnprocessedItems": {}} + + +def test_batch_write_retries_unprocessed_items(): + items = [{"PK": {"S": "a"}}, {"PK": {"S": "b"}}] + first = [{"PutRequest": {"Item": items[1]}}] + client = FakeDdb(unprocessed_first=first) + written = mod._batch_write_all(client, items, sleep=lambda _s: None) + assert written == 2 + assert len(client.calls) == 2 + assert client.calls[1] == first + + +def test_batch_write_raises_if_unprocessed_remain(): + items = [{"PK": {"S": "a"}}] + stuck = [{"PutRequest": {"Item": items[0]}}] + client = FakeDdb(unprocessed_first=stuck) + client._always = True + + def always_unprocessed(RequestItems): + client.calls.append(RequestItems[mod.TABLE]) + return {"UnprocessedItems": {mod.TABLE: stuck}} + + client.batch_write_item = always_unprocessed + try: + mod._batch_write_all(client, items, sleep=lambda _s: None, max_attempts=3) + raise AssertionError("expected RuntimeError") + except RuntimeError as exc: + assert "unprocessed" in str(exc) diff --git a/tests/scripts/test_recreate_holiday_schedules.py b/tests/scripts/test_recreate_holiday_schedules.py new file mode 100644 index 0000000..be7885a --- /dev/null +++ b/tests/scripts/test_recreate_holiday_schedules.py @@ -0,0 +1,57 @@ +"""recreate_holiday_schedules skips past at() expressions, not only EndDate.""" + +from datetime import datetime, timezone +from zoneinfo import ZoneInfo +import importlib.util +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[2] + + +def _load(): + spec = importlib.util.spec_from_file_location( + "recreate_holiday_schedules", + ROOT / "scripts" / "cutover" / "recreate_holiday_schedules.py", + ) + mod = importlib.util.module_from_spec(spec) + sys.modules["recreate_holiday_schedules"] = mod + spec.loader.exec_module(mod) + return mod + + +mod = _load() + + +def test_schedule_when_parses_at_expression_in_eastern(): + detail = { + "ScheduleExpression": "at(2026-07-04T08:00:00)", + "ScheduleExpressionTimezone": "America/New_York", + } + when = mod.schedule_when(detail) + expected = datetime(2026, 7, 4, 8, 0, 0, tzinfo=ZoneInfo("America/New_York")).astimezone( + timezone.utc + ) + assert when == expected + + +def test_past_at_expression_is_before_now_without_end_date(): + detail = { + "ScheduleExpression": "at(2020-01-01T08:00:00)", + "ScheduleExpressionTimezone": "America/New_York", + } + when = mod.schedule_when(detail) + assert when is not None + assert when < datetime.now(timezone.utc) + assert "EndDate" not in detail + assert "StartDate" not in detail + + +def test_future_at_expression_is_kept(): + detail = { + "ScheduleExpression": "at(2099-12-25T17:00:00)", + "ScheduleExpressionTimezone": "America/New_York", + } + when = mod.schedule_when(detail) + assert when is not None + assert when > datetime.now(timezone.utc)