blob: 407044aeba3ecbce342e5403f635d065dc8ef373 [file]
"""Utility functions for the release tool."""
import argparse
import collections.abc
import fnmatch
import os
import pathlib
import re
from packaging.version import parse as parse_version
from dev.release.git import Git
REPO_URL = "https://github.com/bazel-contrib/rules_python"
def semver_type(value):
"""Argparse type validator for semantic versions."""
if not re.match(r"^\d+\.\d+\.\d+(rc\d+)?$", value):
raise argparse.ArgumentTypeError(
f"'{value}' is not a valid semantic version (X.Y.Z or X.Y.ZrcN)"
)
return value
_EXCLUDE_PATTERNS = [
"./.agents/*",
"./.git/*",
"./.github/*",
"./.bazelci/*",
"./.bcr/*",
"./bazel-*/*",
"./CONTRIBUTING.md",
"./RELEASING.md",
"./dev/release/*",
"./tests/tools/private/release/*",
]
def is_excluded_version_placeholder_path(path: pathlib.Path | str) -> bool:
"""Checks if a path matches any version placeholder exclusion patterns."""
path_str = str(path)
if not path_str.startswith("./") and not path_str.startswith("/"):
path_str = f"./{path_str}"
return any(fnmatch.fnmatch(path_str, pattern) for pattern in _EXCLUDE_PATTERNS)
def _iter_version_placeholder_files() -> collections.abc.Iterator[pathlib.Path]:
for root, dirs, files in os.walk(".", topdown=True):
# Filter directories in-place
dirs[:] = [
d
for d in dirs
if not is_excluded_version_placeholder_path(os.path.join(root, d))
]
for filename in files:
filepath = pathlib.Path(root) / filename
if is_excluded_version_placeholder_path(filepath):
continue
yield filepath
def get_latest_version(git=None):
"""Gets the latest version from git tags."""
if git is None:
git = Git(os.getcwd())
tags = git.get_tags()
versions = [
(tag, parse_version(tag))
for tag in tags
if re.match(r"^\d+\.\d+\.\d+(rc\d+)?$", tag.strip())
]
if not versions:
raise RuntimeError("No git tags found matching X.Y.Z or X.Y.ZrcN format.")
versions.sort(key=lambda v: v[1])
latest_tag, latest_version = versions[-1]
if latest_version.is_prerelease:
raise ValueError(f"The latest version is a pre-release version: {latest_tag}")
stable_versions = [tag for tag, version in versions if not version.is_prerelease]
if not stable_versions:
raise ValueError("No stable git tags found matching X.Y.Z format.")
return stable_versions[-1]
def get_latest_rc_tag(version, remote=None, git=None):
"""Queries git tags and returns the highest RC tag for the version."""
if git is None:
git = Git(os.getcwd())
if remote:
tags = git.get_remote_tags(remote)
else:
tags = git.get_tags()
pattern = rf"^{re.escape(version)}-rc\d+$"
rc_tags = [tag.strip() for tag in tags if re.match(pattern, tag.strip())]
if not rc_tags:
return None
rc_tags.sort(key=parse_version)
return rc_tags[-1]
def should_increment_minor():
"""Checks if the minor version should be incremented."""
for filepath in _iter_version_placeholder_files():
try:
with open(filepath, "r") as f:
content = f.read()
except (IOError, UnicodeDecodeError):
continue
if "VERSION_NEXT_FEATURE" in content:
return True
return False
def determine_next_version(branch_name=None, git=None, is_patch=False):
"""Determines the next version based on git tags and the current branch."""
if git is None:
git = Git(os.getcwd())
if branch_name is None:
branch_name = git.get_current_branch()
if branch_name:
release_match = re.match(r"^release/(\d+)\.(\d+)$", branch_name)
if release_match:
branch_major = int(release_match.group(1))
branch_minor = int(release_match.group(2))
print(
f"Detected release branch: {branch_name} (targeting"
f" {branch_major}.{branch_minor}.x)"
)
tags = git.get_tags()
matching_patches = []
for tag in tags:
tag = tag.strip()
m = re.match(rf"^{branch_major}\.{branch_minor}\.(\d+)$", tag)
if m:
matching_patches.append(int(m.group(1)))
if matching_patches:
latest_patch = max(matching_patches)
next_version = f"{branch_major}.{branch_minor}.{latest_patch + 1}"
print(
f"Latest tag on this branch is"
f" {branch_major}.{branch_minor}.{latest_patch}. Next"
f" version: {next_version}"
)
return next_version
else:
next_version = f"{branch_major}.{branch_minor}.0"
print(
f"No stable tags found for {branch_major}.{branch_minor}.x."
f" Next version: {next_version}"
)
return next_version
latest_version = get_latest_version(git=git)
major, minor, patch = [int(n) for n in latest_version.split(".")]
if not is_patch and should_increment_minor():
return f"{major}.{minor + 1}.0"
else:
return f"{major}.{minor}.{patch + 1}"
def replace_version_next_in_files(
filepaths: collections.abc.Iterable[pathlib.Path], version: str
) -> list[pathlib.Path]:
"""Replaces VERSION_NEXT_* placeholders with version in the specified files.
Args:
filepaths: An iterable of pathlib.Path objects to process.
version: The release version string to replace placeholders with.
Returns:
List of pathlib.Path objects for files that were modified.
"""
modified: list[pathlib.Path] = []
for path in filepaths:
if is_excluded_version_placeholder_path(path):
continue
try:
content = path.read_text(encoding="utf-8")
except (IOError, UnicodeDecodeError):
continue
if "VERSION_NEXT_FEATURE" in content or "VERSION_NEXT_PATCH" in content:
new_content = content.replace("VERSION_NEXT_FEATURE", version)
new_content = new_content.replace("VERSION_NEXT_PATCH", version)
path.write_text(new_content, encoding="utf-8")
modified.append(path)
return modified
def replace_version_next(version: str) -> list[pathlib.Path]:
"""Replaces all VERSION_NEXT_* placeholders with the new version.
Args:
version: The release version string to replace placeholders with.
Returns:
List of pathlib.Path objects for files that were modified.
"""
return replace_version_next_in_files(_iter_version_placeholder_files(), version)
def parse_pr_list(value: str) -> list[str]:
"""Parses a comma or space separated list of PR references.
PR references can be numbers (optionally prefixed with '#') or URLs.
"""
if not value:
return []
# Split by space and/or comma
return [p for p in re.split(r"[\s,]+", value.strip()) if p]
def set_github_output(name: str, value: str) -> None:
"""Sets a GitHub Actions output parameter if GITHUB_OUTPUT is set."""
if github_output := os.environ.get("GITHUB_OUTPUT"):
with open(github_output, "a", encoding="utf-8") as f:
f.write(f"{name}={value}\n")
def format_exception(e: BaseException) -> str:
"""Formats an exception to a string, including any attached PEP 678 notes."""
msg = str(e)
notes = getattr(e, "__notes__", None)
if not notes:
return msg
return "\n".join(filter(None, [msg] + [str(note) for note in notes]))