from __future__ import annotations import asyncio from dataclasses import dataclass, field from datetime import timedelta from types import SimpleNamespace import pytest from fastapi.testclient import TestClient import app.rotation.repository as repository_module import app.rotation.service as rotation_module from app.core.errors import ConflictError, ExternalServiceError from app.core.time import from_iso, to_iso, utc_now from app.integrations.aws_client import InstanceSnapshot, StaticIpSnapshot from app.rotation.repository import FleetRun, RotationLeaseLostError, RotationRepository from app.rotation.service import RotationService Timeline = list[tuple[object, ...]] class SimulatedCrash(BaseException): """Represents process death, so RotationService must not handle it as a normal failure.""" @dataclass(slots=True) class StaticResource: name: str ip_address: str region: str is_attached: bool attached_to: str | None @dataclass class FakeLightsailClient: instance_name: str old_ip: str new_ip: str timeline: Timeline region: str = "us-east-1" state: str = "running" reported_is_static_ip: bool | None = None hide_attached_static_ips: bool = False fail_release: bool = False crash_after_detach: bool = False mutations: list[str] = field(default_factory=list) fail_on_action_call: dict[str, int] = field(default_factory=dict) action_call_counts: dict[str, int] = field(default_factory=dict) def __post_init__(self) -> None: self.public_ip = self.old_ip self.old_static_name = f"old-{self.instance_name}" self.resources = { self.old_static_name: StaticResource( name=self.old_static_name, ip_address=self.old_ip, region=self.region, is_attached=True, attached_to=self.instance_name, ) } def get_instance(self) -> InstanceSnapshot: attached = any(resource.is_attached for resource in self.resources.values()) is_static_ip = ( attached if self.reported_is_static_ip is None else self.reported_is_static_ip ) self.timeline.append(("aws:get_instance", self.instance_name, self.public_ip)) return InstanceSnapshot( name=self.instance_name, state=self.state, public_ip=self.public_ip, is_static_ip=is_static_ip, region=self.region, ) def get_attached_static_ips(self, instance_name: str) -> list[StaticIpSnapshot]: assert instance_name == self.instance_name self.timeline.append(("aws:get_attached", self.instance_name)) if self.hide_attached_static_ips: return [] return [ self._snapshot(resource) for resource in self.resources.values() if resource.is_attached and resource.attached_to == instance_name ] def get_static_ip(self, static_ip_name: str | None) -> StaticIpSnapshot | None: self.timeline.append(("aws:get_static", self.instance_name, static_ip_name)) if static_ip_name is None: return None resource = self.resources.get(static_ip_name) return self._snapshot(resource) if resource else None def allocate_static_ip(self, static_ip_name: str) -> None: call_number = self._mutation("allocate", static_ip_name) assert static_ip_name not in self.resources self.resources[static_ip_name] = StaticResource( name=static_ip_name, ip_address=self.new_ip, region=self.region, is_attached=False, attached_to=None, ) self._raise_injected_failure("allocate", call_number) def detach_static_ip(self, static_ip_name: str) -> None: call_number = self._mutation("detach", static_ip_name) resource = self.resources[static_ip_name] assert resource.is_attached resource.is_attached = False resource.attached_to = None self.public_ip = "192.0.2.254" if self.crash_after_detach: self.crash_after_detach = False raise SimulatedCrash("controller crashed after AWS detached the old IP") self._raise_injected_failure("detach", call_number) def attach_static_ip(self, static_ip_name: str, instance_name: str | None = None) -> None: target = instance_name or self.instance_name call_number = self._mutation("attach", static_ip_name) assert target == self.instance_name assert not any(resource.is_attached for resource in self.resources.values()) resource = self.resources[static_ip_name] resource.is_attached = True resource.attached_to = target self.public_ip = resource.ip_address self._raise_injected_failure("attach", call_number) def release_static_ip(self, static_ip_name: str) -> None: call_number = self._mutation("release", static_ip_name) if self.fail_release: raise ExternalServiceError( "release failed", service="aws", code="AWS_RELEASE_FAILED", ) resource = self.resources[static_ip_name] assert not resource.is_attached del self.resources[static_ip_name] self._raise_injected_failure("release", call_number) def _mutation(self, action: str, static_ip_name: str) -> int: self.mutations.append(action) self.timeline.append((f"aws:{action}", self.instance_name, static_ip_name)) call_number = self.action_call_counts.get(action, 0) + 1 self.action_call_counts[action] = call_number return call_number def _raise_injected_failure(self, action: str, call_number: int) -> None: if self.fail_on_action_call.get(action) != call_number: return del self.fail_on_action_call[action] raise ExternalServiceError( f"injected {action} failure", service="aws", code=f"INJECTED_{action.upper()}_FAILURE", ) @staticmethod def _snapshot(resource: StaticResource) -> StaticIpSnapshot: return StaticIpSnapshot( name=resource.name, ip_address=resource.ip_address, is_attached=resource.is_attached, attached_to=resource.attached_to, arn=f"arn:aws:lightsail:::StaticIp/{resource.name}", region=resource.region, ) class FakeCloudflareClient: def __init__(self, record_name: str, initial_ip: str, timeline: Timeline) -> None: self.record_name = record_name self.zone_id = f"zone-{record_name}" self.timeline = timeline self.record: dict[str, object] = { "id": f"record-{record_name}", "name": record_name, "type": "A", "content": initial_ip, "proxied": False, "ttl": 60, } self.fail_update_once = False def get_a_record(self) -> dict[str, object]: self.timeline.append(("cf:get", self.record_name, self.record["content"])) return dict(self.record) def upsert_a_record( self, ip: str, *, expected_current: str | None = None, ) -> dict[str, object]: current = str(self.record["content"]) self.timeline.append(("cf:update", self.record_name, ip, expected_current)) assert expected_current is None or current in {expected_current, ip} if self.fail_update_once: self.fail_update_once = False raise ExternalServiceError( "injected DNS update failure", service="cloudflare", code="INJECTED_DNS_UPDATE_FAILURE", ) self.record["content"] = ip self.record["proxied"] = False return dict(self.record) class FakeIntegrationService: def ensure_configured(self) -> None: return None def credentials(self) -> tuple[SimpleNamespace, dict[str, str | None]]: return ( SimpleNamespace(use_default_aws_credentials=False), { "aws_access_key_id": "test-access-key", "aws_secret_access_key": "test-secret-key", "aws_session_token": None, "cloudflare_api_token": "test-cloudflare-token", }, ) class RotationHarness: def __init__( self, client: TestClient, monkeypatch: pytest.MonkeyPatch, ) -> None: self.container = client.app.state.container self.monkeypatch = monkeypatch self.timeline: Timeline = [] self.aws_by_instance: dict[str, FakeLightsailClient] = {} self.cloudflare_by_record: dict[str, FakeCloudflareClient] = {} self.probe_failures: dict[str, int] = {} self.clock = utc_now() self.integration_service = FakeIntegrationService() def lightsail_factory(**kwargs: object) -> FakeLightsailClient: return self.aws_by_instance[str(kwargs["instance_name"])] def cloudflare_factory(**kwargs: object) -> FakeCloudflareClient: return self.cloudflare_by_record[str(kwargs["record_name"])] monkeypatch.setattr( rotation_module.LightsailClient, "from_credentials", staticmethod(lightsail_factory), ) monkeypatch.setattr(rotation_module, "CloudflareClient", cloudflare_factory) monkeypatch.setattr(rotation_module, "utc_now", lambda: self.clock) self.service = self.new_service() @property def repository(self) -> RotationRepository: return self.container.rotation_repository def new_service(self) -> RotationService: service = RotationService( self.container.rotation_repository, self.integration_service, # type: ignore[arg-type] poll_interval_seconds=0, ) async def probe_socks5(ip: str, port: int) -> None: self.timeline.append(("socks", ip, port)) failures = self.probe_failures.get(ip, 0) if failures: self.probe_failures[ip] = failures - 1 raise ExternalServiceError( "injected SOCKS failure", service="socks5", code="INJECTED_SOCKS_FAILURE", ) async def advance_grace(run_id: str) -> None: service._require_lease(run_id) self.timeline.append(("grace_wait", run_id)) self.clock += timedelta(seconds=61) self.monkeypatch.setattr(service, "_probe_socks5", probe_socks5) self.monkeypatch.setattr(service, "_sleep_with_lease", advance_grace) return service def add_instance( self, instance_id: str, *, old_ip: str, new_ip: str, state: str = "running", reported_is_static_ip: bool | None = None, hide_attached_static_ips: bool = False, fail_release: bool = False, crash_after_detach: bool = False, ) -> str: instance_name = f"node-{instance_id}" record_name = f"{instance_id}.example.com" self.container.fleet_repository.create_instance( { "id": instance_id, "display_name": f"Proxy {instance_id}", "aws_region": "us-east-1", "lightsail_instance_name": instance_name, "cloudflare_zone_name": "example.com", "cloudflare_zone_id": f"zone-{instance_id}", "cloudflare_record_name": record_name, "socks_port": 1080, "proxy_health_check": True, "health_timeout_seconds": 10, "release_grace_seconds": 60, "enabled": True, } ) self.aws_by_instance[instance_name] = FakeLightsailClient( instance_name=instance_name, old_ip=old_ip, new_ip=new_ip, timeline=self.timeline, state=state, reported_is_static_ip=reported_is_static_ip, hide_attached_static_ips=hide_attached_static_ips, fail_release=fail_release, crash_after_detach=crash_after_detach, ) self.cloudflare_by_record[record_name] = FakeCloudflareClient( record_name, old_ip, self.timeline, ) return instance_id def create_group(self, group_id: str, member_ids: list[str]) -> str: self.container.fleet_repository.create_group( { "id": group_id, "name": f"Group {group_id}", "enabled": True, "interval_minutes": 15, "member_ids": member_ids, } ) return group_id @pytest.fixture def harness(client: TestClient, monkeypatch: pytest.MonkeyPatch) -> RotationHarness: return RotationHarness(client, monkeypatch) def mutating_actions(timeline: Timeline) -> list[tuple[object, ...]]: return [ event for event in timeline if event[0] in { "aws:allocate", "aws:detach", "aws:attach", "aws:release", "cf:update", } ] async def drive_run_until_pause_or_terminal( harness: RotationHarness, run_id: str, ) -> FleetRun: for _ in range(8): retry = await harness.service._execute(run_id) current = harness.repository.get(run_id) assert current is not None if current.status not in {"queued", "running"} or not retry: return current raise AssertionError("rotation did not reach a stable state") @pytest.mark.asyncio async def test_static_ip_rotation_releases_old_ip_only_after_all_safety_gates( harness: RotationHarness, ) -> None: instance_id = harness.add_instance( "one", old_ip="198.51.100.10", new_ip="198.51.100.20", ) run = harness.repository.create_for_instance(instance_id) assert await harness.service._execute(run.id) is False completed = harness.repository.get(run.id) item = harness.repository.list_items(run.id)[0] aws = harness.aws_by_instance["node-one"] assert completed is not None assert completed.status == "succeeded" assert completed.succeeded_items == 1 assert item.status == "succeeded" assert item.stage == "succeeded" assert item.old_static_ip_name == aws.old_static_name assert aws.old_static_name not in aws.resources assert len(aws.resources) == 1 assert next(iter(aws.resources.values())).ip_address == "198.51.100.20" assert next(iter(aws.resources.values())).is_attached is True action_names = [str(event[0]) for event in harness.timeline] attach_index = action_names.index("aws:attach") first_socks_index = action_names.index("socks") update_index = action_names.index("cf:update") verified_index = next( index for index, event in enumerate(harness.timeline) if index > update_index and event[:3] == ("cf:get", "one.example.com", "198.51.100.20") ) grace_index = action_names.index("grace_wait") last_socks_index = len(action_names) - 1 - action_names[::-1].index("socks") release_index = action_names.index("aws:release") assert attach_index < first_socks_index < update_index assert update_index < verified_index < grace_index < last_socks_index < release_index assert [event[0] for event in mutating_actions(harness.timeline)] == [ "aws:allocate", "aws:detach", "aws:attach", "cf:update", "aws:release", ] expected_stages = [ "queued", "preflight", "new_allocating", "old_detaching", "new_attaching", "health_checking", "dns_updating", "dns_verifying", "old_releasing", ] stages = [event["stage"] for event in harness.repository.list_events(run.id)] assert [stage for stage in stages if stage in expected_stages] == expected_stages @pytest.mark.asyncio async def test_release_failure_enters_cleanup_pending_and_resume_is_idempotent( harness: RotationHarness, monkeypatch: pytest.MonkeyPatch, ) -> None: instance_id = harness.add_instance( "one", old_ip="198.51.100.10", new_ip="198.51.100.20", fail_release=True, ) run = harness.repository.create_for_instance(instance_id) await harness.service._execute(run.id) pending = harness.repository.get(run.id) item = harness.repository.list_items(run.id)[0] aws = harness.aws_by_instance["node-one"] assert pending is not None assert pending.status == "cleanup_pending" assert pending.active_slot == 1 assert pending.error_code == "AWS_RELEASE_FAILED" assert item.status == "cleanup_pending" assert item.stage == "old_releasing" assert aws.resources[aws.old_static_name].is_attached is False assert harness.cloudflare_by_record["one.example.com"].record["content"] == item.new_ip aws.fail_release = False monkeypatch.setattr(harness.service, "_spawn", lambda _run_id: None) resumed = harness.service.resume(run.id) assert resumed.status == "running" assert resumed.trigger == "recovery" assert harness.repository.list_items(run.id)[0].attempt_count == 1 await harness.service._execute(run.id) completed = harness.repository.get(run.id) assert completed is not None assert completed.status == "succeeded" assert aws.old_static_name not in aws.resources assert aws.mutations.count("allocate") == 1 assert aws.mutations.count("detach") == 1 assert aws.mutations.count("attach") == 1 assert aws.mutations.count("release") == 2 @pytest.mark.asyncio @pytest.mark.parametrize( ("failure_stage", "expected_code", "must_detach_new"), [ ("old_detaching", "INJECTED_DETACH_FAILURE", False), ("new_attaching", "INJECTED_ATTACH_FAILURE", True), ("health_checking", "INJECTED_SOCKS_FAILURE", True), ("dns_updating", "INJECTED_DNS_UPDATE_FAILURE", True), ], ) async def test_automatic_rollback_restores_old_route_before_releasing_new_static_ip( harness: RotationHarness, failure_stage: str, expected_code: str, must_detach_new: bool, ) -> None: old_ip = "198.51.100.10" new_ip = "198.51.100.20" instance_id = harness.add_instance("one", old_ip=old_ip, new_ip=new_ip) aws = harness.aws_by_instance["node-one"] cloudflare = harness.cloudflare_by_record["one.example.com"] if failure_stage == "old_detaching": aws.fail_on_action_call["detach"] = 1 elif failure_stage == "new_attaching": aws.fail_on_action_call["attach"] = 1 elif failure_stage == "health_checking": harness.probe_failures[new_ip] = 1 else: cloudflare.fail_update_once = True run = harness.repository.create_for_instance(instance_id) assert await harness.service._execute(run.id) is True rollback_checkpoint = harness.repository.list_items(run.id)[0] assert rollback_checkpoint.status == "running" assert rollback_checkpoint.stage == "rollback_reconciling" assert rollback_checkpoint.rollback_from_stage == failure_stage assert rollback_checkpoint.rollback_reason_code == expected_code terminal = await drive_run_until_pause_or_terminal(harness, run.id) item = harness.repository.list_items(run.id)[0] assert terminal.status == "failed" assert terminal.error_code == expected_code assert terminal.active_slot is None assert item.status == "failed" assert item.stage == "rollback_new_releasing" assert item.rollback_from_stage == failure_stage assert item.rollback_reason_code == expected_code assert cloudflare.record["content"] == old_ip assert aws.public_ip == old_ip assert aws.old_static_name in aws.resources assert aws.resources[aws.old_static_name].is_attached is True assert aws.resources[aws.old_static_name].attached_to == "node-one" assert item.new_static_ip_name not in aws.resources old_attach_index = next( index for index, event in enumerate(harness.timeline) if event == ("aws:attach", "node-one", aws.old_static_name) ) release_new_index = next( index for index, event in enumerate(harness.timeline) if event == ("aws:release", "node-one", item.new_static_ip_name) ) restored_dns_index = next( index for index, event in enumerate(harness.timeline) if index > old_attach_index and event == ("cf:get", "one.example.com", old_ip) ) restored_socks_index = next( index for index, event in enumerate(harness.timeline) if index > old_attach_index and event == ("socks", old_ip, 1080) ) assert old_attach_index < restored_dns_index < restored_socks_index < release_new_index new_detach_indices = [ index for index, event in enumerate(harness.timeline) if event == ("aws:detach", "node-one", item.new_static_ip_name) ] if must_detach_new: assert len(new_detach_indices) == 1 assert new_detach_indices[0] < old_attach_index else: assert new_detach_indices == [] assert not any( event == ("aws:release", "node-one", aws.old_static_name) for event in harness.timeline ) @pytest.mark.asyncio async def test_rollback_failure_pauses_and_resume_continues_idempotently( harness: RotationHarness, monkeypatch: pytest.MonkeyPatch, ) -> None: old_ip = "198.51.100.10" new_ip = "198.51.100.20" instance_id = harness.add_instance("one", old_ip=old_ip, new_ip=new_ip) aws = harness.aws_by_instance["node-one"] harness.probe_failures[new_ip] = 1 aws.fail_on_action_call["detach"] = 2 run = harness.repository.create_for_instance(instance_id) paused = await drive_run_until_pause_or_terminal(harness, run.id) paused_item = harness.repository.list_items(run.id)[0] assert paused.status == "needs_attention" assert paused.error_code == "INJECTED_DETACH_FAILURE" assert paused_item.status == "needs_attention" assert paused_item.stage == "rollback_new_detaching" assert paused_item.rollback_from_stage == "health_checking" assert paused_item.rollback_reason_code == "INJECTED_SOCKS_FAILURE" assert aws.resources[str(paused_item.new_static_ip_name)].is_attached is False monkeypatch.setattr(harness.service, "_spawn", lambda _run_id: None) resumed = harness.service.resume(run.id) assert resumed.status == "running" assert harness.repository.list_items(run.id)[0].attempt_count == 1 terminal = await drive_run_until_pause_or_terminal(harness, run.id) completed_item = harness.repository.list_items(run.id)[0] assert terminal.status == "failed" assert terminal.error_code == "INJECTED_SOCKS_FAILURE" assert completed_item.status == "failed" assert completed_item.rollback_reason_code == "INJECTED_SOCKS_FAILURE" assert aws.mutations.count("allocate") == 1 assert aws.mutations.count("detach") == 2 assert aws.mutations.count("attach") == 2 assert aws.mutations.count("release") == 1 assert aws.old_static_name in aws.resources assert aws.resources[aws.old_static_name].is_attached is True assert completed_item.new_static_ip_name not in aws.resources @pytest.mark.asyncio async def test_preflight_failure_has_no_cloud_side_effects(harness: RotationHarness) -> None: instance_id = harness.add_instance( "stopped", old_ip="198.51.100.10", new_ip="198.51.100.20", state="stopped", ) run = harness.repository.create_for_instance(instance_id) await harness.service._execute(run.id) failed = harness.repository.get(run.id) item = harness.repository.list_items(run.id)[0] assert failed is not None assert failed.status == "failed" assert failed.error_code == "INSTANCE_NOT_RUNNING" assert item.status == "failed" assert item.stage == "preflight" assert mutating_actions(harness.timeline) == [] assert harness.repository.get_active() is None @pytest.mark.asyncio async def test_preflight_rejects_static_flag_without_attached_resource( harness: RotationHarness, ) -> None: instance_id = harness.add_instance( "inconsistent", old_ip="198.51.100.10", new_ip="198.51.100.20", reported_is_static_ip=True, hide_attached_static_ips=True, ) run = harness.repository.create_for_instance(instance_id) await harness.service._execute(run.id) failed = harness.repository.get(run.id) item = harness.repository.list_items(run.id)[0] assert failed is not None assert failed.status == "failed" assert failed.error_code == "STATIC_IP_STATE_CONFLICT" assert item.stage == "preflight" assert mutating_actions(harness.timeline) == [] @pytest.mark.asyncio async def test_group_rotates_instances_strictly_in_member_order( harness: RotationHarness, ) -> None: first = harness.add_instance( "one", old_ip="198.51.100.10", new_ip="198.51.100.20", ) second = harness.add_instance( "two", old_ip="203.0.113.10", new_ip="203.0.113.20", ) group_id = harness.create_group("primary", [first, second]) run = harness.repository.create_for_group(group_id) await harness.service._execute(run.id) completed = harness.repository.get(run.id) items = harness.repository.list_items(run.id) assert completed is not None assert completed.status == "succeeded" assert completed.succeeded_items == 2 assert [item.instance_id for item in items] == [first, second] assert [item.status for item in items] == ["succeeded", "succeeded"] first_release = next( index for index, event in enumerate(harness.timeline) if event[0] == "aws:release" and event[1] == "node-one" ) first_second_instance_call = next( index for index, event in enumerate(harness.timeline) if str(event[0]).startswith("aws:") and event[1] == "node-two" ) second_allocate = next( index for index, event in enumerate(harness.timeline) if event[0] == "aws:allocate" and event[1] == "node-two" ) assert first_release < first_second_instance_call <= second_allocate @pytest.mark.asyncio async def test_group_stops_before_next_member_when_cleanup_is_pending( harness: RotationHarness, ) -> None: first = harness.add_instance( "one", old_ip="198.51.100.10", new_ip="198.51.100.20", fail_release=True, ) second = harness.add_instance( "two", old_ip="203.0.113.10", new_ip="203.0.113.20", ) group_id = harness.create_group("primary", [first, second]) run = harness.repository.create_for_group(group_id) await harness.service._execute(run.id) pending = harness.repository.get(run.id) items = harness.repository.list_items(run.id) assert pending is not None assert pending.status == "cleanup_pending" assert [item.status for item in items] == ["cleanup_pending", "queued"] assert not any( str(event[0]).startswith("aws:") and event[1] == "node-two" for event in harness.timeline ) @pytest.mark.asyncio async def test_terminal_group_failure_cancels_all_unstarted_members( harness: RotationHarness, ) -> None: first = harness.add_instance( "one", old_ip="198.51.100.10", new_ip="198.51.100.20", ) second = harness.add_instance( "two", old_ip="203.0.113.10", new_ip="203.0.113.20", ) harness.probe_failures["198.51.100.20"] = 1 group_id = harness.create_group("primary", [first, second]) run = harness.repository.create_for_group(group_id) terminal = await drive_run_until_pause_or_terminal(harness, run.id) items = harness.repository.list_items(run.id) assert terminal.status == "failed" assert [item.status for item in items] == ["failed", "cancelled"] assert items[1].stage == "cancelled" assert items[1].error_code == "BATCH_ABORTED" assert items[1].finished_at is not None assert not any( str(event[0]).startswith("aws:") and event[1] == "node-two" for event in harness.timeline ) @pytest.mark.asyncio @pytest.mark.parametrize("terminal_status", ["succeeded", "failed"]) async def test_scheduled_group_rebases_next_run_from_terminal_time( harness: RotationHarness, monkeypatch: pytest.MonkeyPatch, terminal_status: str, ) -> None: instance_id = harness.add_instance( "one", old_ip="198.51.100.10", new_ip="198.51.100.20", ) if terminal_status == "failed": harness.probe_failures["198.51.100.20"] = 1 group_id = harness.create_group("scheduled", [instance_id]) initial_group = harness.container.fleet_repository.get_group(group_id) assert initial_group is not None assert initial_group.next_run_at is not None initial_next_run = from_iso(initial_group.next_run_at) run = harness.repository.create_for_group(group_id, trigger="scheduled") in_progress_group = harness.container.fleet_repository.get_group(group_id) assert in_progress_group is not None assert in_progress_group.next_run_at is None terminal_time = from_iso(run.started_at) + timedelta(hours=2) monkeypatch.setattr(repository_module, "utc_now", lambda: terminal_time) monkeypatch.setattr( repository_module, "to_iso", lambda value=None: to_iso(value or terminal_time), ) terminal = await drive_run_until_pause_or_terminal(harness, run.id) updated_group = harness.container.fleet_repository.get_group(group_id) assert terminal.status == terminal_status assert terminal.finished_at is not None assert updated_group is not None assert updated_group.last_run_at == terminal.finished_at assert updated_group.next_run_at is not None assert initial_next_run < from_iso(terminal.finished_at) assert ( from_iso(updated_group.next_run_at) - from_iso(terminal.finished_at) == timedelta(minutes=15) ) def test_lease_owner_is_fenced_after_another_worker_takes_over( harness: RotationHarness, ) -> None: instance_id = harness.add_instance( "one", old_ip="198.51.100.10", new_ip="198.51.100.20", ) run = harness.repository.create_for_instance(instance_id) item = harness.repository.list_items(run.id)[0] assert harness.repository.acquire_lease(run.id, "worker-one") is True harness.repository.transition_run( run.id, status="running", expected_owner="worker-one", ) harness.repository.transition_item( run.id, item.id, status="running", expected_owner="worker-one", ) with harness.repository.database.connect() as connection, connection: connection.execute( "UPDATE fleet_runs SET lease_until = ? WHERE id = ?", (to_iso(utc_now() - timedelta(seconds=1)), run.id), ) assert harness.repository.acquire_lease(run.id, "worker-two") is True with pytest.raises(RotationLeaseLostError, match="ROTATION_LEASE_LOST"): harness.repository.transition_item( run.id, item.id, stage="new_allocating", expected_owner="worker-one", ) with pytest.raises(RotationLeaseLostError, match="ROTATION_LEASE_LOST"): harness.repository.add_event( run.id, item_id=item.id, stage="preflight", message="stale worker event", expected_owner="worker-one", ) updated = harness.repository.transition_item( run.id, item.id, stage="new_allocating", expected_owner="worker-two", ) assert updated.stage == "new_allocating" assert not any( event["message"] == "stale worker event" for event in harness.repository.list_events(run.id) ) @pytest.mark.asyncio async def test_lease_heartbeat_fences_worker_before_any_new_side_effect( harness: RotationHarness, monkeypatch: pytest.MonkeyPatch, ) -> None: renew_calls: list[tuple[str, str]] = [] async def no_wait(_seconds: float) -> None: return None def reject_renewal(run_id: str, owner: str) -> bool: renew_calls.append((run_id, owner)) return False monkeypatch.setattr(rotation_module.asyncio, "sleep", no_wait) monkeypatch.setattr(harness.repository, "renew_lease", reject_renewal) await harness.service._lease_heartbeat("run-with-lost-lease") assert "run-with-lost-lease" in harness.service._lease_lost_runs assert renew_calls == [("run-with-lost-lease", harness.service.worker_id)] with pytest.raises(RotationLeaseLostError, match="ROTATION_LEASE_LOST"): harness.service._require_lease("run-with-lost-lease") assert renew_calls == [("run-with-lost-lease", harness.service.worker_id)] @pytest.mark.asyncio async def test_recover_continues_after_crash_between_cloud_mutation_and_checkpoint( harness: RotationHarness, ) -> None: instance_id = harness.add_instance( "one", old_ip="198.51.100.10", new_ip="198.51.100.20", crash_after_detach=True, ) run = harness.repository.create_for_instance(instance_id) with pytest.raises(SimulatedCrash, match="controller crashed"): await harness.service._execute(run.id) interrupted = harness.repository.get(run.id) item = harness.repository.list_items(run.id)[0] aws = harness.aws_by_instance["node-one"] assert interrupted is not None assert interrupted.status == "running" assert item.stage == "old_detaching" assert aws.resources[aws.old_static_name].is_attached is False with harness.repository.database.connect() as connection, connection: connection.execute( "UPDATE fleet_runs SET lease_until = ? WHERE id = ?", (to_iso(utc_now() - timedelta(seconds=1)), run.id), ) restarted = harness.new_service() restarted.recover() recovery_task = restarted._tasks[run.id] await recovery_task await asyncio.sleep(0) completed = harness.repository.get(run.id) assert completed is not None assert completed.status == "succeeded" assert aws.mutations.count("allocate") == 1 assert aws.mutations.count("detach") == 1 assert aws.mutations.count("attach") == 1 assert aws.mutations.count("release") == 1 assert any( event["stage"] == "recovery" for event in harness.repository.list_events(run.id) ) @pytest.mark.asyncio async def test_cleanup_pending_blocks_other_runs_until_cancelled( harness: RotationHarness, monkeypatch: pytest.MonkeyPatch, ) -> None: first = harness.add_instance( "one", old_ip="198.51.100.10", new_ip="198.51.100.20", fail_release=True, ) second = harness.add_instance( "two", old_ip="203.0.113.10", new_ip="203.0.113.20", ) run = harness.repository.create_for_instance(first) await harness.service._execute(run.id) assert harness.repository.get(run.id).status == "cleanup_pending" # type: ignore[union-attr] monkeypatch.setattr(harness.service, "_spawn", lambda _run_id: None) with pytest.raises(ConflictError) as exc_info: harness.service.start_instance(second) assert exc_info.value.code == "ROTATION_IN_PROGRESS" cloud_events_before_cancel = list(harness.timeline) cancelled = harness.service.cancel(run.id) assert cancelled.status == "cancelled" assert harness.timeline == cloud_events_before_cancel assert harness.repository.get_active() is None next_run = harness.service.start_instance(second) assert next_run.status == "queued" assert next_run.active_slot == 1