import unittest import hashlib import os import base64 import sys from pathlib import Path from unittest.mock import patch, AsyncMock from unittest.mock import MagicMock EMMA_DIR = Path(__file__).resolve().parents[1] / "emma" if str(EMMA_DIR) not in sys.path: sys.path.insert(0, str(EMMA_DIR)) import server # Følger eksisterende mønster KMS_KEY_NAME = "projects/p/locations/l/keyRings/k/cryptoKeys/k" class TestProposeGiteaChange(unittest.IsolatedAsyncioTestCase): @patch('server.handle_get_file_content', new_callable=AsyncMock) async def test_get_file_wrapper_enforces_server_repo(self, mock_handler): sentinel_response = { "requested_ref": "main", "resolved_commit_sha": "a" * 40, "content": "file content" } mock_handler.return_value = sentinel_response params = { "path": "README.md", "ref": "main", "repo": "untrusted/repo", } result = await server.get_file(params) mock_handler.assert_awaited_once_with(params, server.GITEA_REPO) self.assertIs(result, sentinel_response) def setUp(self): """Set up valid parameters for tests.""" self.valid_params = { "repo": server.GITEA_REPO, "branch": "feature/new-idea", "path": "docs/new-file.md", "new_content": "This is new content.", "base_sha": "a" * 40, "commit_message": "A valid commit message.", } def tearDown(self): patch.stopall() @patch('server.create_gitea_change_plan', new_callable=AsyncMock) @patch('server._get_gitea_file_details', new_callable=AsyncMock) @patch('server.resolve_branch_to_commit_sha', new_callable=AsyncMock) async def test_rejects_main_branch_before_io(self, mock_resolve_sha, mock_get_details, mock_create_plan): params = self.valid_params | {"branch": "main"} with self.assertRaisesRegex(ValueError, "Direct writes to protected branch 'main' are not allowed."): await server.propose_gitea_change(params) mock_resolve_sha.assert_not_awaited() mock_get_details.assert_not_awaited() mock_create_plan.assert_not_awaited() @patch('server.create_gitea_change_plan', new_callable=AsyncMock) @patch('server._get_gitea_file_details', new_callable=AsyncMock) @patch('server.resolve_branch_to_commit_sha', new_callable=AsyncMock) async def test_rejects_master_branch_before_io(self, mock_resolve_sha, mock_get_details, mock_create_plan): params = self.valid_params | {"branch": "master"} with self.assertRaisesRegex(ValueError, "Direct writes to protected branch 'master' are not allowed."): await server.propose_gitea_change(params) mock_resolve_sha.assert_not_awaited() mock_get_details.assert_not_awaited() mock_create_plan.assert_not_awaited() @patch("server.create_gitea_change_plan", new_callable=AsyncMock) @patch("server._get_gitea_file_details", new_callable=AsyncMock) @patch("server.resolve_branch_to_commit_sha", new_callable=AsyncMock) @patch("server._validate_admin_gitea_path") async def test_rejects_dot_git_path_via_new_write_policy( self, mock_admin_path_validator, mock_resolve_sha, mock_get_details, mock_create_plan, ): params = self.valid_params | {"path": "some/dir/.git/config"} with self.assertRaisesRegex( ValueError, r"Changes within a '.git' directory are not allowed.", ): await server.propose_gitea_change(params) mock_admin_path_validator.assert_called_once_with(params["path"]) mock_resolve_sha.assert_not_awaited() mock_get_details.assert_not_awaited() mock_create_plan.assert_not_awaited() @patch('server.create_gitea_change_plan', new_callable=AsyncMock) @patch('server._get_gitea_file_details', new_callable=AsyncMock) @patch('server.resolve_branch_to_commit_sha', new_callable=AsyncMock) async def test_rejects_base_sha_mismatch(self, mock_resolve_sha, mock_get_details, mock_create_plan): mock_resolve_sha.return_value = "c" * 40 # Mismatched SHA with self.assertRaisesRegex(ValueError, "Branch head does not match the supplied base_sha"): await server.propose_gitea_change(self.valid_params) mock_resolve_sha.assert_awaited_once() mock_get_details.assert_not_awaited() mock_create_plan.assert_not_awaited() @patch('server.create_gitea_change_plan', new_callable=AsyncMock) @patch('server._get_gitea_file_details', new_callable=AsyncMock) @patch('server.resolve_branch_to_commit_sha', new_callable=AsyncMock) async def test_rejects_no_op_diff(self, mock_resolve_sha, mock_get_details, mock_create_plan): mock_resolve_sha.return_value = self.valid_params["base_sha"] mock_get_details.return_value = (self.valid_params["new_content"], "b" * 40) with self.assertRaisesRegex(ValueError, "Proposed content produces no file change."): await server.propose_gitea_change(self.valid_params) mock_resolve_sha.assert_awaited_once() mock_get_details.assert_awaited_once() mock_create_plan.assert_not_awaited() @patch('server.create_gitea_change_plan', new_callable=AsyncMock) @patch("server.kms_v1.KeyManagementServiceAsyncClient") @patch('server._get_gitea_file_details', new_callable=AsyncMock, return_value=("old content", "b" * 40)) @patch('server.resolve_branch_to_commit_sha', new_callable=AsyncMock) async def test_happy_path_creates_pending_plan(self, mock_resolve_sha, mock_get_details, mock_kms_constructor, mock_create_plan): mock_kms_client = AsyncMock() mock_kms_client.encrypt.return_value = MagicMock(ciphertext=b"encrypted-data") mock_kms_constructor.return_value.__aenter__ = AsyncMock(return_value=mock_kms_client) mock_kms_constructor.return_value.__aexit__ = AsyncMock(return_value=False) mock_resolve_sha.return_value = self.valid_params["base_sha"] with patch.dict(os.environ, {"GITEA_PLAN_KMS_KEY_NAME": KMS_KEY_NAME}): result = await server.propose_gitea_change(self.valid_params) self.assertEqual(result["status"], "PENDING") self.assertEqual(result["repo"], self.valid_params["repo"]) self.assertEqual(result["branch"], self.valid_params["branch"]) self.assertEqual(result["path"], self.valid_params["path"]) self.assertIn("approval_subject_hash", result) expected_hash = hashlib.sha256(self.valid_params["new_content"].encode("utf-8")).hexdigest() self.assertEqual(result["content_hash"], expected_hash) mock_kms_client.encrypt.assert_awaited_once() mock_create_plan.assert_awaited_once() # Verify the object passed to create_gitea_change_plan is the real Pydantic model call_args = mock_create_plan.call_args[0][0] self.assertIsInstance(call_args, server.GiteaChangePlan) self.assertEqual(call_args.repo, self.valid_params["repo"]) self.assertEqual(call_args.branch, self.valid_params["branch"]) self.assertEqual(call_args.path, self.valid_params["path"]) self.assertEqual(call_args.base_sha, self.valid_params["base_sha"]) self.assertEqual(call_args.existing_file_sha, "b" * 40) @patch('server.create_gitea_change_plan', new_callable=AsyncMock) @patch("server.kms_v1.KeyManagementServiceAsyncClient") @patch('server._get_gitea_file_details', new_callable=AsyncMock) @patch('server.resolve_branch_to_commit_sha', new_callable=AsyncMock) async def test_patch2a_missing_kms_key_fails_after_gitea_reads( self, mock_resolve_sha, mock_get_details, mock_kms_constructor, mock_create_plan ): mock_resolve_sha.return_value = self.valid_params["base_sha"] mock_get_details.return_value = ("old content", "b" * 40) with patch.dict(os.environ, {"GITEA_PLAN_KMS_KEY_NAME": ""}): with self.assertRaisesRegex(RuntimeError, "KMS key is not configured"): await server.propose_gitea_change(self.valid_params) mock_resolve_sha.assert_awaited_once() mock_get_details.assert_awaited_once() mock_kms_constructor.assert_not_called() mock_create_plan.assert_not_awaited() @patch('server.create_gitea_change_plan', new_callable=AsyncMock) @patch("server.kms_v1.KeyManagementServiceAsyncClient") @patch('server._get_gitea_file_details', new_callable=AsyncMock) @patch('server.resolve_branch_to_commit_sha', new_callable=AsyncMock) async def test_patch2a_invalid_kms_key_fails_before_kms_construction( self, mock_resolve_sha, mock_get_details, mock_kms_constructor, mock_create_plan ): mock_resolve_sha.return_value = self.valid_params["base_sha"] mock_get_details.return_value = ("old content", "b" * 40) with patch.dict(os.environ, {"GITEA_PLAN_KMS_KEY_NAME": "invalid-key-format"}): with self.assertRaisesRegex(RuntimeError, "invalid format"): await server.propose_gitea_change(self.valid_params) mock_resolve_sha.assert_awaited_once() mock_get_details.assert_awaited_once() mock_kms_constructor.assert_not_called() mock_create_plan.assert_not_awaited() @patch('server.create_gitea_change_plan', new_callable=AsyncMock) @patch("server.kms_v1.KeyManagementServiceAsyncClient") @patch('server._get_gitea_file_details', new_callable=AsyncMock, return_value=("old content", "b" * 40)) @patch('server.resolve_branch_to_commit_sha', new_callable=AsyncMock) async def test_patch2a_empty_kms_ciphertext_is_rejected( self, mock_resolve_sha, mock_get_details, mock_kms_constructor, mock_create_plan ): mock_resolve_sha.return_value = self.valid_params["base_sha"] mock_kms_client = AsyncMock() mock_kms_client.encrypt.return_value = MagicMock(ciphertext=b"") mock_kms_constructor.return_value.__aenter__ = AsyncMock(return_value=mock_kms_client) mock_kms_constructor.return_value.__aexit__ = AsyncMock(return_value=False) with patch.dict(os.environ, {"GITEA_PLAN_KMS_KEY_NAME": KMS_KEY_NAME}): with self.assertRaisesRegex(RuntimeError, "invalid empty ciphertext"): await server.propose_gitea_change(self.valid_params) mock_kms_client.encrypt.assert_awaited_once_with( request={ "name": KMS_KEY_NAME, "plaintext": self.valid_params["new_content"].encode("utf-8"), } ) mock_create_plan.assert_not_awaited() @patch('server.create_gitea_change_plan', new_callable=AsyncMock) @patch("server.kms_v1.KeyManagementServiceAsyncClient") @patch('server._get_gitea_file_details', new_callable=AsyncMock, return_value=("old content", "b" * 40)) @patch('server.resolve_branch_to_commit_sha', new_callable=AsyncMock) async def test_patch2a_non_bytes_kms_ciphertext_is_rejected( self, mock_resolve_sha, mock_get_details, mock_kms_constructor, mock_create_plan ): mock_resolve_sha.return_value = self.valid_params["base_sha"] mock_kms_client = AsyncMock() mock_kms_client.encrypt.return_value = MagicMock(ciphertext=None) mock_kms_constructor.return_value.__aenter__ = AsyncMock(return_value=mock_kms_client) mock_kms_constructor.return_value.__aexit__ = AsyncMock(return_value=False) with patch.dict(os.environ, {"GITEA_PLAN_KMS_KEY_NAME": KMS_KEY_NAME}): with self.assertRaisesRegex(RuntimeError, "invalid empty ciphertext"): await server.propose_gitea_change(self.valid_params) mock_kms_client.encrypt.assert_awaited_once_with( request={ "name": KMS_KEY_NAME, "plaintext": self.valid_params["new_content"].encode("utf-8"), } ) mock_create_plan.assert_not_awaited() @patch('server.create_gitea_change_plan', new_callable=AsyncMock) @patch("server.kms_v1.KeyManagementServiceAsyncClient") @patch('server._get_gitea_file_details', new_callable=AsyncMock, return_value=("old content", "b" * 40)) @patch('server.resolve_branch_to_commit_sha', new_callable=AsyncMock) async def test_patch2a_successful_encrypted_proposal( self, mock_resolve_sha, mock_get_details, mock_kms_constructor, mock_create_plan ): mock_resolve_sha.return_value = self.valid_params["base_sha"] mock_kms_client = AsyncMock() mock_kms_client.encrypt.return_value = MagicMock(ciphertext=b"encrypted-data") mock_kms_constructor.return_value.__aenter__ = AsyncMock(return_value=mock_kms_client) mock_kms_constructor.return_value.__aexit__ = AsyncMock(return_value=False) with patch.dict(os.environ, {"GITEA_PLAN_KMS_KEY_NAME": KMS_KEY_NAME}): result = await server.propose_gitea_change(self.valid_params) mock_resolve_sha.assert_awaited_once() mock_get_details.assert_awaited_once() mock_kms_client.encrypt.assert_awaited_once_with( request={ "name": KMS_KEY_NAME, "plaintext": self.valid_params["new_content"].encode("utf-8"), } ) mock_create_plan.assert_awaited_once() persisted_plan = mock_create_plan.call_args.args[0] self.assertIsInstance(persisted_plan, server.GiteaChangePlan) self.assertIsInstance( persisted_plan.encrypted_payload, server.EncryptedPayloadV1, ) self.assertEqual(persisted_plan.encrypted_payload.payload_version, "1") self.assertEqual(persisted_plan.encrypted_payload.kms_key_name, KMS_KEY_NAME) self.assertEqual( persisted_plan.encrypted_payload.ciphertext, base64.b64encode(b"encrypted-data").decode("ascii"), ) self.assertNotIn("new_content", persisted_plan.model_dump(mode="json")) expected_content_hash = hashlib.sha256(self.valid_params["new_content"].encode("utf-8")).hexdigest() self.assertEqual(persisted_plan.content_hash, expected_content_hash) self.assertEqual(result["content_hash"], expected_content_hash) self.assertEqual( persisted_plan.approval_subject_hash, server._calculate_approval_subject_hash(persisted_plan), ) self.assertEqual(result["approval_subject_hash"], persisted_plan.approval_subject_hash) for field in ("new_content", "encrypted_payload", "ciphertext", "kms_key_name"): self.assertNotIn(field, result) def test_legacy_plan_parses_without_patch2a_fields(self): current_plan = server.GiteaChangePlan( repo="o/r", branch="b", path="p", base_sha="a" * 40, content_hash="c" * 64, commit_message="m", unified_diff="d" ) legacy_dict = current_plan.model_dump(mode="json") legacy_dict.pop("encrypted_payload", None) legacy_dict.pop("approval_subject_hash", None) legacy_plan = server.GiteaChangePlan(**legacy_dict) self.assertIsNone(legacy_plan.encrypted_payload) self.assertIsNone(legacy_plan.approval_subject_hash) def test_malformed_base64_in_hash_raises_error(self): plan_with_bad_payload = server.GiteaChangePlan( plan_id="plan-123", repo="o/r", branch="b", path="p", base_sha="a"*40, content_hash="c"*64, commit_message="m", unified_diff="d", encrypted_payload=server.EncryptedPayloadV1( kms_key_name=KMS_KEY_NAME, ciphertext="not-valid-base64!" ) ) with self.assertRaisesRegex(ValueError, "Encrypted payload ciphertext is invalid"): server._calculate_approval_subject_hash(plan_with_bad_payload) def test_approval_subject_hash_sensitivity(self): base_plan = server.GiteaChangePlan( plan_id="plan-123", repo="o/r", branch="b", path="p", base_sha="a" * 40, existing_file_sha="b" * 40, content_hash="c" * 64, commit_message="m", unified_diff="d", encrypted_payload=server.EncryptedPayloadV1( kms_key_name=KMS_KEY_NAME, ciphertext=base64.b64encode(b"secret").decode("ascii"), ), ) base_hash = server._calculate_approval_subject_hash(base_plan) fields_to_test = { "plan_id": "plan-456", "repo": "o/r2", "branch": "b2", "path": "p2", "base_sha": "A" * 40, "existing_file_sha": "B" * 40, "content_hash": "C" * 64, "commit_message": "m2", "unified_diff": "d2", } for field, new_value in fields_to_test.items(): with self.subTest(field=field): modified_plan = base_plan.model_copy(update={field: new_value}) new_hash = server._calculate_approval_subject_hash(modified_plan) self.assertNotEqual(base_hash, new_hash) with self.subTest(field="kms_key_name"): modified_plan_kms = base_plan.model_copy(deep=True) modified_plan_kms.encrypted_payload.kms_key_name = "projects/p/locations/l/keyRings/k/cryptoKeys/k2" self.assertNotEqual(base_hash, server._calculate_approval_subject_hash(modified_plan_kms)) with self.subTest(field="ciphertext"): modified_plan_cipher = base_plan.model_copy(deep=True) modified_plan_cipher.encrypted_payload.ciphertext = base64.b64encode(b"secret2").decode("ascii") self.assertNotEqual(base_hash, server._calculate_approval_subject_hash(modified_plan_cipher)) class TestIsGiteaChangePlanExpired(unittest.TestCase): def _create_test_plan(self, expires_at): return server.GiteaChangePlan( repo="o/r", branch="b", path="p", base_sha="a" * 40, content_hash="c" * 64, commit_message="m", unified_diff="d", expires_at=expires_at, ) def test_future_aware_is_not_expired(self): from datetime import datetime, timedelta, timezone future_time = datetime.now(timezone.utc) + timedelta(days=1) plan = self._create_test_plan(future_time) self.assertFalse(server.is_gitea_change_plan_expired(plan)) def test_past_aware_is_expired(self): from datetime import datetime, timedelta, timezone past_time = datetime.now(timezone.utc) - timedelta(days=1) plan = self._create_test_plan(past_time) self.assertTrue(server.is_gitea_change_plan_expired(plan)) def test_past_naive_is_expired(self): from datetime import datetime, timedelta past_time_naive = datetime.utcnow() - timedelta(days=1) plan = self._create_test_plan(past_time_naive) self.assertTrue(server.is_gitea_change_plan_expired(plan)) def test_default_status_is_pending_constant(self): from datetime import datetime, timedelta, timezone future_time = datetime.now(timezone.utc) + timedelta(days=1) plan = self._create_test_plan(future_time) self.assertEqual( plan.status, server.GITEA_CHANGE_PLAN_STATUS_PENDING, ) self.assertEqual(server.GITEA_CHANGE_PLAN_STATUS_PENDING, "PENDING") self.assertEqual(server.GITEA_CHANGE_PLAN_STATUS_APPROVED, "APPROVED") self.assertEqual(server.GITEA_CHANGE_PLAN_STATUS_APPLYING, "APPLYING") self.assertEqual(server.GITEA_CHANGE_PLAN_STATUS_APPLIED, "APPLIED") self.assertEqual(server.GITEA_CHANGE_PLAN_STATUS_REJECTED, "REJECTED") self.assertEqual(server.GITEA_CHANGE_PLAN_STATUS_EXPIRED, "EXPIRED") class TestGiteaChangePlanStatusTransitions(unittest.TestCase): def test_allowed_transitions(self): self.assertTrue( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_PENDING, server.GITEA_CHANGE_PLAN_STATUS_APPROVED, ) ) self.assertTrue( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_PENDING, server.GITEA_CHANGE_PLAN_STATUS_REJECTED, ) ) self.assertTrue( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_PENDING, server.GITEA_CHANGE_PLAN_STATUS_EXPIRED, ) ) self.assertTrue( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_APPROVED, server.GITEA_CHANGE_PLAN_STATUS_APPLYING, ) ) self.assertTrue( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_APPROVED, server.GITEA_CHANGE_PLAN_STATUS_EXPIRED, ) ) self.assertTrue( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_APPLYING, server.GITEA_CHANGE_PLAN_STATUS_APPLIED, ) ) def test_disallowed_transitions(self): self.assertFalse( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_PENDING, server.GITEA_CHANGE_PLAN_STATUS_PENDING, ) ) self.assertFalse( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_APPROVED, server.GITEA_CHANGE_PLAN_STATUS_APPLIED, ) ) self.assertFalse( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_APPLYING, server.GITEA_CHANGE_PLAN_STATUS_APPROVED, ) ) self.assertFalse( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_APPLIED, server.GITEA_CHANGE_PLAN_STATUS_APPLYING, ) ) self.assertFalse( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_REJECTED, server.GITEA_CHANGE_PLAN_STATUS_APPROVED, ) ) self.assertFalse( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_EXPIRED, server.GITEA_CHANGE_PLAN_STATUS_APPROVED, ) ) self.assertFalse( server.is_valid_gitea_change_plan_status_transition( "UNKNOWN", server.GITEA_CHANGE_PLAN_STATUS_PENDING ) ) self.assertFalse( server.is_valid_gitea_change_plan_status_transition( server.GITEA_CHANGE_PLAN_STATUS_PENDING, "UNKNOWN" ) ) class TestTransitionGiteaChangePlanStatus(unittest.IsolatedAsyncioTestCase): @patch("google.cloud.firestore.AsyncClient") async def test_rejects_invalid_transition_before_firestore(self, mock_db_client): with self.assertRaisesRegex( ValueError, r"^Invalid Gitea change plan status transition$" ): await server.transition_gitea_change_plan_status( plan_id="plan-123", expected_status=server.GITEA_CHANGE_PLAN_STATUS_PENDING, next_status=server.GITEA_CHANGE_PLAN_STATUS_APPLIED, ) mock_db_client.assert_not_called() @patch("google.cloud.firestore.AsyncClient") async def test_rejects_status_in_updates_before_firestore(self, mock_db_client): with self.assertRaisesRegex( ValueError, r"^Gitea change plan updates cannot include status$" ): await server.transition_gitea_change_plan_status( plan_id="plan-123", expected_status=server.GITEA_CHANGE_PLAN_STATUS_PENDING, next_status=server.GITEA_CHANGE_PLAN_STATUS_APPROVED, updates={"status": server.GITEA_CHANGE_PLAN_STATUS_REJECTED}, ) mock_db_client.assert_not_called() @patch("server.get_gitea_change_plan", new_callable=AsyncMock) @patch("google.cloud.firestore.async_transactional") @patch("google.cloud.firestore.AsyncClient") async def test_transitions_matching_status_and_returns_refetched_plan( self, mock_db_client, mock_transactional, mock_get_plan ): def transactional_side_effect(callback): async def wrapped(transaction): return await callback(transaction) return wrapped mock_transactional.side_effect = transactional_side_effect mock_transaction = MagicMock() mock_transaction.update = MagicMock() mock_db_client.return_value.transaction.return_value = mock_transaction snapshot = MagicMock() snapshot.exists = True snapshot.get.return_value = server.GITEA_CHANGE_PLAN_STATUS_PENDING mock_plan_ref = MagicMock() mock_plan_ref.get = AsyncMock(return_value=snapshot) mock_db_client.return_value.collection.return_value.document.return_value = mock_plan_ref final_plan = server.GiteaChangePlan( status=server.GITEA_CHANGE_PLAN_STATUS_APPROVED, repo="o/r", branch="b", path="p", base_sha="a"*40, content_hash="c"*64, commit_message="m", unified_diff="d", approved_by="admin@example.com" ) mock_get_plan.return_value = final_plan updates = {"approved_by": "admin@example.com"} result = await server.transition_gitea_change_plan_status( plan_id="plan-123", expected_status=server.GITEA_CHANGE_PLAN_STATUS_PENDING, next_status=server.GITEA_CHANGE_PLAN_STATUS_APPROVED, updates=updates, ) mock_transaction.update.assert_called_once_with( mock_plan_ref, { "approved_by": "admin@example.com", "status": server.GITEA_CHANGE_PLAN_STATUS_APPROVED, }, ) self.assertEqual(updates, {"approved_by": "admin@example.com"}) mock_get_plan.assert_awaited_once_with("plan-123") self.assertIs(result, final_plan) @patch("server.get_gitea_change_plan", new_callable=AsyncMock) @patch("google.cloud.firestore.async_transactional") @patch("google.cloud.firestore.AsyncClient") async def test_rejects_stale_status_inside_transaction( self, mock_db_client, mock_transactional, mock_get_plan ): def transactional_side_effect(callback): async def wrapped(transaction): return await callback(transaction) return wrapped mock_transactional.side_effect = transactional_side_effect mock_transaction = MagicMock() mock_transaction.update = MagicMock() mock_db_client.return_value.transaction.return_value = mock_transaction snapshot = MagicMock() snapshot.exists = True snapshot.get.return_value = server.GITEA_CHANGE_PLAN_STATUS_APPROVED mock_plan_ref = MagicMock() mock_plan_ref.get = AsyncMock(return_value=snapshot) mock_db_client.return_value.collection.return_value.document.return_value = mock_plan_ref with self.assertRaisesRegex(ValueError, r"^Gitea change plan status changed$"): await server.transition_gitea_change_plan_status( plan_id="plan-123", expected_status=server.GITEA_CHANGE_PLAN_STATUS_PENDING, next_status=server.GITEA_CHANGE_PLAN_STATUS_APPROVED, ) mock_transaction.update.assert_not_called() mock_get_plan.assert_not_awaited() if __name__ == "__main__": unittest.main()