diff --git a/src/scripts/new_proposal.py b/src/scripts/new_proposal.py
index d0ef228dde51..0ad491f5267a 100755
--- a/src/scripts/new_proposal.py
+++ b/src/scripts/new_proposal.py
@@ -8,6 +8,7 @@ Exceptions. See /LICENSE for license information.
SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
"""
+import argparse
import os
import re
import shlex
@@ -15,12 +16,6 @@ import shutil
import subprocess
import sys
-_USAGE = """Usage:
- ./new-proposal.py
[]
-
-Generates a branch and PR for a new proposal with the specified title.
-"""
-
_PROMPT = """This will:
- Create and switch to a new branch named '%s'.
- Create a new proposal titled '%s'.
@@ -29,7 +24,46 @@ _PROMPT = """This will:
Continue? (Y/n) """
-def _FillTemplate(template_path, title, pr_num):
+def _exit(error):
+ """Wraps sys.exit for testing."""
+ sys.exit(error)
+
+
+def _parse_args(args=None):
+ """Parses command-line arguments and flags."""
+ parser = argparse.ArgumentParser(
+ description="Generates a branch and PR for a new proposal with the "
+ "specified title."
+ )
+ parser.add_argument(
+ "title", metavar="TITLE", help="The title of the proposal.",
+ )
+ parser.add_argument(
+ "--branch",
+ metavar="BRANCH",
+ help="The name of the branch. Automatically generated from the title "
+ "by default.",
+ )
+ return parser.parse_args(args=args)
+
+
+def _calculate_branch(parsed_args):
+ """Returns the branch name."""
+ if parsed_args.branch:
+ return parsed_args.branch
+ # Only use the first 20 chars of the title for branch names.
+ return "proposal-%s" % (parsed_args.title.lower().replace(" ", "-")[0:20])
+
+
+def _find_tool(tool):
+ """Checks if a tool is present."""
+ tool_path = shutil.which(tool)
+ if not tool_path:
+ _exit("ERROR: Missing the '%s' command-line tool." % tool)
+ return tool_path
+
+
+def _fill_template(template_path, title, pr_num):
"""Fills out template TODO fields."""
with open(template_path) as template_file:
content = template_file.read()
@@ -43,16 +77,16 @@ def _FillTemplate(template_path, title, pr_num):
return content
-def _Run(argv, check=True):
+def _run(argv, check=True):
"""Runs a command."""
cmd = " ".join([shlex.quote(x) for x in argv])
print("\n+ RUNNING: %s" % cmd, file=sys.stderr)
p = subprocess.run(argv)
if check and p.returncode != 0:
- sys.exit("ERROR: Command failed: %s" % cmd)
+ _exit("ERROR: Command failed: %s" % cmd)
-def _RunPRCreate(argv):
+def _run_pr_create(argv):
"""Runs a command and returns the PR#."""
cmd = " ".join([shlex.quote(x) for x in argv])
print("\n+ RUNNING: %s" % cmd, file=sys.stderr)
@@ -61,34 +95,24 @@ def _RunPRCreate(argv):
out = out.decode("utf-8")
print(out, end="")
if p.returncode != 0:
- sys.exit("ERROR: Command failed: %s" % cmd)
+ _exit("ERROR: Command failed: %s" % cmd)
match = re.search(
r"^https://github.com/[^/]+/[^/]+/pull/(\d+)$", out, re.MULTILINE
)
if not match:
- sys.exit("ERROR: Failed to find PR# in output.")
+ _exit("ERROR: Failed to find PR# in output.")
return int(match[1])
-if __name__ == "__main__":
- # Require an argument.
- if len(sys.argv) not in (2, 3):
- sys.exit(_USAGE)
- title = sys.argv[1]
- branch = None
- if len(sys.argv) == 3:
- branch = sys.argv[2]
+def main():
+ parsed_args = _parse_args()
+ title = parsed_args.title
+ branch = _calculate_branch(parsed_args)
- # Verify git and gh are available.
- git_bin = shutil.which("git")
- if not git_bin:
- sys.exit("ERROR: Missing `git` CLI.")
- gh_bin = shutil.which("gh")
- if not gh_bin:
- sys.exit("ERROR: Missing `gh` CLI.")
- precommit_bin = shutil.which("pre-commit")
- if not precommit_bin:
- sys.exit("ERROR: Missing `pre-commit` CLI.")
+ # Verify tools are available.
+ git_bin = _find_tool("git")
+ gh_bin = _find_tool("gh")
+ precommit_bin = _find_tool("pre-commit")
# Ensure a good working directory.
proposals_dir = os.path.realpath(
@@ -99,33 +123,29 @@ if __name__ == "__main__":
# Verify there are no uncommitted changes.
p = subprocess.run([git_bin, "diff-index", "--quiet", "HEAD", "--"])
if p.returncode != 0:
- sys.exit("ERROR: There are uncommitted changes in your git repo.")
-
- # Only use the first 20 chars of the title for branch names.
- if not branch:
- branch = "proposal-%s" % (title.lower().replace(" ", "-")[0:20])
+ _exit("ERROR: There are uncommitted changes in your git repo.")
# Prompt before proceeding.
response = "?"
while response not in ("y", "n", ""):
response = input(_PROMPT % (branch, title)).lower()
if response == "n":
- sys.exit("ERROR: Cancelled")
+ _exit("ERROR: Cancelled")
# Create a proposal branch.
- _Run([git_bin, "checkout", "-b", branch, "trunk"])
- _Run([git_bin, "push", "-u", "origin", branch])
+ _run([git_bin, "checkout", "-b", branch, "trunk"])
+ _run([git_bin, "push", "-u", "origin", branch])
# Copy template.md to a temp file.
template_path = os.path.join(proposals_dir, "template.md")
temp_path = os.path.join(proposals_dir, "new-proposal.tmp")
shutil.copyfile(template_path, temp_path)
- _Run([git_bin, "add", temp_path])
- _Run([git_bin, "commit", "-m", "Creating new proposal: %s" % title])
+ _run([git_bin, "add", temp_path])
+ _run([git_bin, "commit", "-m", "Creating new proposal: %s" % title])
# Create a PR with WIP+proposal labels.
- _Run([git_bin, "push"])
- pr_num = _RunPRCreate(
+ _run([git_bin, "push"])
+ pr_num = _run_pr_create(
[
gh_bin,
"pr",
@@ -142,18 +162,30 @@ if __name__ == "__main__":
# Remove the temp file, create p####.md, and fill in PR information.
os.remove(temp_path)
final_path = os.path.join(proposals_dir, "p%04d.md" % pr_num)
- content = _FillTemplate(template_path, title, pr_num)
+ content = _fill_template(template_path, title, pr_num)
with open(final_path, "w") as final_file:
final_file.write(content)
- _Run([git_bin, "add", temp_path, final_path])
- _Run([precommit_bin, "run"], check=False) # Needs a ToC update.
- _Run([git_bin, "add", final_path, os.path.join(proposals_dir, "README.md")])
- _Run([git_bin, "commit", "-m", "Filling out template with PR %d" % pr_num])
+ _run([git_bin, "add", temp_path, final_path])
+ _run([precommit_bin, "run"], check=False) # Needs a ToC update.
+ _run([git_bin, "add", final_path, os.path.join(proposals_dir, "README.md")])
+ _run(
+ [
+ git_bin,
+ "commit",
+ "--amend",
+ "-m",
+ "Filling out template with PR %d" % pr_num,
+ ]
+ )
# Push the PR update.
- _Run([git_bin, "push"])
+ _run([git_bin, "push", "--force-with-lease"])
print(
"\nCreated PR %d for %s. Make changes to:\n %s"
% (pr_num, title, final_path)
)
+
+
+if __name__ == "__main__":
+ main()
diff --git a/src/scripts/new_proposal_test.py b/src/scripts/new_proposal_test.py
new file mode 100644
index 000000000000..495d01a24fca
--- /dev/null
+++ b/src/scripts/new_proposal_test.py
@@ -0,0 +1,67 @@
+#!/usr/bin/env python3
+
+"""Tests for new_proposal.py."""
+
+__copyright__ = """
+Part of the Carbon Language project, under the Apache License v2.0 with LLVM
+Exceptions. See /LICENSE for license information.
+SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+"""
+
+import unittest
+from unittest import mock
+
+import new_proposal
+
+
+class FakeExitError(Exception):
+ pass
+
+
+def _fake_exit(message):
+ raise FakeExitError(message)
+
+
+class TestNewProposal(unittest.TestCase):
+ def test_calculate_branch_short(self):
+ parsed_args = new_proposal._parse_args(["foo bar"])
+ self.assertEqual(
+ new_proposal._calculate_branch(parsed_args), "proposal-foo-bar"
+ )
+
+ def test_calculate_branch_long(self):
+ parsed_args = new_proposal._parse_args(
+ ["A really long long long title"]
+ )
+ self.assertEqual(
+ new_proposal._calculate_branch(parsed_args),
+ "proposal-a-really-long-long-l",
+ )
+
+ def test_calculate_branch_flag(self):
+ parsed_args = new_proposal._parse_args(["--branch=wiz", "foo"])
+ self.assertEqual(new_proposal._calculate_branch(parsed_args), "wiz")
+
+ def test_fill_template(self):
+ content = new_proposal._fill_template(
+ "../../proposals/template.md", "TITLE", 123
+ )
+ self.assertTrue(content.startswith("# TITLE\n\n"), content)
+ self.assertTrue(
+ "[Pull request](https://github.com/carbon-language/carbon-lang/"
+ "pull/123)" in content,
+ content,
+ )
+
+ def test_run_success(self):
+ new_proposal._run(["true"])
+
+ def test_run_failure(self):
+ with mock.patch(
+ "new_proposal._exit", side_effect=_fake_exit
+ ) as mock_exit:
+ self.assertRaises(FakeExitError, new_proposal._run, ["false"])
+
+
+if __name__ == "__main__":
+ unittest.main()