From 3bf57483518623ec3bb583ac2fea5bb553fc2bcf Mon Sep 17 00:00:00 2001 From: Richard Lundeen Date: Thu, 10 Sep 2026 18:01:03 -0700 Subject: [PATCH] FIX: Allow reviewed migration rewrite temporarily Pin the temporary Alembic immutability exception to the exact blob reviewed in #2583. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../enforce_alembic_revision_immutability.py | 23 +++++++-- ...t_enforce_alembic_revision_immutability.py | 50 ++++++++++++++++++- 2 files changed, 69 insertions(+), 4 deletions(-) diff --git a/build_scripts/enforce_alembic_revision_immutability.py b/build_scripts/enforce_alembic_revision_immutability.py index e7891e352c..ed5521144c 100644 --- a/build_scripts/enforce_alembic_revision_immutability.py +++ b/build_scripts/enforce_alembic_revision_immutability.py @@ -16,6 +16,9 @@ _VERSIONS_PATH = "pyrit/memory/alembic/versions/" _MERGE_QUEUE_REF_PREFIX = "refs/heads/gh-readonly-queue/" +# Temporary: remove this exception immediately after #2583 merges. +_APPROVED_REWRITE_PATH = f"{_VERSIONS_PATH}1b3d5f7a9c2e_persist_scored_expectation.py" +_APPROVED_REWRITE_BLOB = "3cb40abf25ffafe4cbd1ddfda213e02e54458d91" def _git(*args: str) -> subprocess.CompletedProcess[str]: @@ -26,10 +29,23 @@ def _git_stdout(*args: str) -> str: return _git(*args).stdout.strip() +def _find_violations(*, output: str, diff_spec: list[str]) -> list[str]: + changes = [line for line in output.splitlines() if line and not line.startswith("A")] + return [change for change in changes if not _is_approved_rewrite(change=change, diff_spec=diff_spec)] + + +def _is_approved_rewrite(*, change: str, diff_spec: list[str]) -> bool: + if change != f"M\t{_APPROVED_REWRITE_PATH}": + return False + + object_spec = f":{_APPROVED_REWRITE_PATH}" if diff_spec == ["--cached"] else f"HEAD:{_APPROVED_REWRITE_PATH}" + return _git_stdout("rev-parse", "--verify", object_spec) == _APPROVED_REWRITE_BLOB + + def _get_violations(diff_spec: list[str]) -> list[str]: """Return lines from ``git diff --name-status`` that are not pure additions.""" output = _git_stdout("diff", "--name-status", *diff_spec, "--", _VERSIONS_PATH) - return [line for line in output.splitlines() if line and not line.startswith("A")] + return _find_violations(output=output, diff_spec=diff_spec) def _in_ci() -> bool: @@ -91,9 +107,10 @@ def has_revision_violations() -> bool: # automatically, so we don't need a separate merge-base call. When # the base is missing (shallow clone) git exits non-zero. base = f"origin/{base_ref}" if base_ref else "origin/main" - pr_diff = _git("diff", "--name-status", f"{base}...HEAD", "--", _VERSIONS_PATH) + diff_spec = [f"{base}...HEAD"] + pr_diff = _git("diff", "--name-status", *diff_spec, "--", _VERSIONS_PATH) if pr_diff.returncode == 0: - violations = [line for line in pr_diff.stdout.strip().splitlines() if line and not line.startswith("A")] + violations = _find_violations(output=pr_diff.stdout, diff_spec=diff_spec) if violations: _report(violations) return True diff --git a/tests/unit/build_scripts/test_enforce_alembic_revision_immutability.py b/tests/unit/build_scripts/test_enforce_alembic_revision_immutability.py index 752f31354a..b9cd80c722 100644 --- a/tests/unit/build_scripts/test_enforce_alembic_revision_immutability.py +++ b/tests/unit/build_scripts/test_enforce_alembic_revision_immutability.py @@ -9,17 +9,65 @@ import pytest from build_scripts.enforce_alembic_revision_immutability import ( + _APPROVED_REWRITE_BLOB, + _APPROVED_REWRITE_PATH, + _find_violations, _on_release_branch, has_revision_violations, ) MODIFIED_REVISION = "M\tpyrit/memory/alembic/versions/b2f4c6a8d1e3_add_conversations_table.py" +APPROVED_REWRITE = f"M\t{_APPROVED_REWRITE_PATH}" def _completed(stdout: str = "", returncode: int = 0) -> subprocess.CompletedProcess: return subprocess.CompletedProcess(args=["git"], returncode=returncode, stdout=stdout, stderr="") +@pytest.mark.parametrize( + "diff_spec, expected_object", + [ + pytest.param(["--cached"], f":{_APPROVED_REWRITE_PATH}", id="staged"), + pytest.param(["origin/main...HEAD"], f"HEAD:{_APPROVED_REWRITE_PATH}", id="pull-request"), + pytest.param(["HEAD~1..HEAD"], f"HEAD:{_APPROVED_REWRITE_PATH}", id="post-merge"), + ], +) +def test_approved_rewrite_passes_each_check(diff_spec: list[str], expected_object: str) -> None: + def _git_result(*args: str) -> subprocess.CompletedProcess: + if args[:2] == ("diff", "--name-status"): + stdout = APPROVED_REWRITE if diff_spec[0] in args else "" + return _completed(stdout=stdout) + if args == ("rev-parse", "--verify", expected_object): + return _completed(stdout=_APPROVED_REWRITE_BLOB) + return _completed() + + environment = {"GITHUB_REF": "refs/heads/main", "GITHUB_BASE_REF": ""} + with patch.dict(os.environ, environment, clear=True): + with patch( + "build_scripts.enforce_alembic_revision_immutability._git", + side_effect=_git_result, + ) as mock_git: + assert has_revision_violations() is False + + assert _APPROVED_REWRITE_BLOB == "3cb40abf25ffafe4cbd1ddfda213e02e54458d91" + assert any(call.args == ("rev-parse", "--verify", expected_object) for call in mock_git.call_args_list) + + +def test_approved_rewrite_with_different_content_fails() -> None: + with patch( + "build_scripts.enforce_alembic_revision_immutability._git_stdout", + return_value="0000000000000000000000000000000000000000", + ): + assert _find_violations(output=APPROVED_REWRITE, diff_spec=["--cached"]) == [APPROVED_REWRITE] + + +def test_different_revision_still_fails() -> None: + with patch("build_scripts.enforce_alembic_revision_immutability._git_stdout") as mock_git_stdout: + assert _find_violations(output=MODIFIED_REVISION, diff_spec=["HEAD~1..HEAD"]) == [MODIFIED_REVISION] + + mock_git_stdout.assert_not_called() + + @pytest.mark.parametrize( "environment, expected", [ @@ -107,7 +155,7 @@ def test_release_pull_request_reports_modified_revision() -> None: def _modified(*args, **kwargs): if "diff" in args and any(arg == "origin/releases/v1.1.0...HEAD" for arg in args): - return SimpleNamespace(returncode=0, stdout=f"M\t{MODIFIED_REVISION}\n", stderr="") + return SimpleNamespace(returncode=0, stdout=f"{MODIFIED_REVISION}\n", stderr="") return SimpleNamespace(returncode=0, stdout="", stderr="") with patch.dict(os.environ, {"GITHUB_BASE_REF": "releases/v1.1.0"}, clear=False):