mirror of
https://github.com/Sea-Haven-Industries/afterhours-shift-manager.git
synced 2026-09-30 05:33:12 +00:00
fix(cutover): retry DDB unprocessed items and skip past at() holidays
Unprocessed BatchWriteItem rows and leftover past at() schedules would drop roster data or abort holiday recreation during prod cutover.
This commit is contained in:
parent
af5e2183e0
commit
634a454d50
4 changed files with 182 additions and 23 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
64
tests/scripts/test_copy_dynamodb.py
Normal file
64
tests/scripts/test_copy_dynamodb.py
Normal file
|
|
@ -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)
|
||||
57
tests/scripts/test_recreate_holiday_schedules.py
Normal file
57
tests/scripts/test_recreate_holiday_schedules.py
Normal file
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue