Some checks are pending
Check Python Version Consistency / Check Python Version (push) Waiting to run
616 lines
27 KiB
Python
616 lines
27 KiB
Python
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()
|