"""Tests for graceful MCP maintenance-drain mode (#659). Acceptance coverage: 1. Enter/exit is durable and audited (DB substrate). 2. New work assignment stops during drain (allocator WAIT). 3. Mutations deferred except allowlisted safety ops. 4. Sessions can observe drain state. 5. Fail-closed on unreadable drain state. """ from __future__ import annotations import sys import tempfile import unittest from pathlib import Path from unittest import mock sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) import maintenance_drain from control_plane_db import ControlPlaneDB from allocator_service import WorkCandidate, allocate_next_work, OUTCOME_WAIT class TestDrainDecisions(unittest.TestCase): def test_inactive_allows_mutations_and_assignment(self): decision = maintenance_drain.classify_mutation("create_pr", None) self.assertTrue(decision["allowed"]) self.assertFalse(decision["deferred"]) assign = maintenance_drain.classify_assignment(None) self.assertTrue(assign["assignment_allowed"]) def test_draining_defers_non_allowlisted_mutation(self): record = {"state": "draining", "remote": "prgs", "org": "o", "repo": "r"} decision = maintenance_drain.classify_mutation("create_pr", record) self.assertFalse(decision["allowed"]) self.assertTrue(decision["deferred"]) self.assertEqual(decision["reason_code"], maintenance_drain.BLOCKER_DRAIN_ACTIVE) self.assertIn("create_pr", decision["reasons"][0]) def test_allowlisted_safety_ops_pass_during_drain(self): record = {"state": "draining"} for task in ( "heartbeat_issue_lock", "gitea_release_reviewer_pr_lease", "write_session_checkpoint", "exit_maintenance_drain", ): with self.subTest(task=task): decision = maintenance_drain.classify_mutation(task, record) self.assertTrue(decision["allowed"], decision) def test_assignment_stopped_during_drain(self): record = {"state": "draining", "reason": "upgrade"} decision = maintenance_drain.classify_assignment(record) self.assertFalse(decision["assignment_allowed"]) self.assertEqual( decision["reason_code"], maintenance_drain.REASON_ASSIGNMENT_STOPPED ) def test_unknown_state_fails_closed(self): with self.assertRaises(maintenance_drain.MaintenanceDrainError): maintenance_drain.normalize_state("drainig") def test_status_payload_always_answers(self): inactive = maintenance_drain.status_payload(None, remote="prgs", org="o", repo="r") self.assertFalse(inactive["draining"]) self.assertTrue(inactive["reads_permitted"]) active = maintenance_drain.status_payload( {"state": "draining", "reason": "reboot", "requested_by": "ops"}, remote="prgs", org="o", repo="r", ) self.assertTrue(active["draining"]) self.assertTrue(active["assignment_stopped"]) self.assertTrue(active["mutations_deferred"]) self.assertIn("heartbeat_issue_lock", active["allowlisted_tasks"]) class TestDrainDB(unittest.TestCase): def setUp(self): self._tmp = tempfile.TemporaryDirectory() self.addCleanup(self._tmp.cleanup) self.db = ControlPlaneDB(db_path=str(Path(self._tmp.name) / "cp.sqlite3")) def test_enter_exit_idempotent_and_audited(self): first = self.db.set_maintenance_drain( remote="prgs", org="Scaled-Tech-Consulting", repo="Gitea-Tools", state="draining", reason="planned restart", requested_by="sysadmin", requested_by_profile="prgs-controller", session_id="s1", ) self.assertTrue(first["transitioned"]) self.assertEqual(first["state"], "draining") self.assertTrue(maintenance_drain.is_draining(first["record"])) again = self.db.set_maintenance_drain( remote="prgs", org="Scaled-Tech-Consulting", repo="Gitea-Tools", state="draining", reason="still draining", requested_by="sysadmin", requested_by_profile="prgs-controller", session_id="s1", ) self.assertFalse(again["transitioned"]) self.assertEqual(again["record"]["entered_at"], first["record"]["entered_at"]) exited = self.db.set_maintenance_drain( remote="prgs", org="Scaled-Tech-Consulting", repo="Gitea-Tools", state="inactive", reason="done", requested_by="sysadmin", requested_by_profile="prgs-controller", session_id="s1", ) self.assertTrue(exited["transitioned"]) self.assertFalse(maintenance_drain.is_draining(exited["record"])) self.assertTrue(exited["record"]["exited_at"]) # Events recorded for transitions only (enter + exit). with self.db._tx(immediate=False) as conn: rows = conn.execute( "SELECT event_type FROM events WHERE event_type LIKE 'maintenance_drain_%' " "ORDER BY event_id" ).fetchall() types = [r[0] for r in rows] self.assertEqual(types, ["maintenance_drain_enter", "maintenance_drain_exit"]) def test_read_missing_is_none_not_error(self): self.assertIsNone( self.db.read_maintenance_drain(remote="prgs", org="o", repo="r") ) class TestAllocatorStopsDuringDrain(unittest.TestCase): def setUp(self): self._tmp = tempfile.TemporaryDirectory() self.addCleanup(self._tmp.cleanup) self.db = ControlPlaneDB(db_path=str(Path(self._tmp.name) / "cp.sqlite3")) def test_allocate_returns_wait_while_draining(self): self.db.set_maintenance_drain( remote="prgs", org="Scaled-Tech-Consulting", repo="Gitea-Tools", state="draining", reason="test", requested_by="tester", ) candidates = [ WorkCandidate( kind="issue", number=659, title="drain", labels=("status:ready",), priority=20, ) ] result = allocate_next_work( self.db, role="author", session_id="test-session", remote="prgs", org="Scaled-Tech-Consulting", repo="Gitea-Tools", apply=False, candidates=candidates, username="jcwalker3", profile_name="prgs-author", ) self.assertEqual(result["outcome"], OUTCOME_WAIT) self.assertIsNone(result.get("selected")) self.assertEqual( result.get("reason_code"), maintenance_drain.REASON_ASSIGNMENT_STOPPED, ) self.assertTrue(result["maintenance_drain"]["draining"]) def test_allocate_works_when_inactive(self): candidates = [ WorkCandidate( kind="issue", number=659, title="drain", labels=("status:ready",), priority=20, ) ] result = allocate_next_work( self.db, role="author", session_id="test-session-2", remote="prgs", org="Scaled-Tech-Consulting", repo="Gitea-Tools", apply=False, candidates=candidates, username="jcwalker3", profile_name="prgs-author", ) self.assertNotEqual( result.get("reason_code"), maintenance_drain.REASON_ASSIGNMENT_STOPPED, ) class TestCapabilityMap(unittest.TestCase): def test_drain_tasks_mapped(self): import task_capability_map as tcm self.assertEqual( tcm.required_permission("enter_maintenance_drain"), "runtime.maintenance_drain", ) self.assertEqual( tcm.required_permission("exit_maintenance_drain"), "runtime.maintenance_drain", ) self.assertEqual( tcm.required_permission("maintenance_drain_status"), "gitea.read", ) if __name__ == "__main__": unittest.main()