diff --git a/rhodecode/__init__.py b/rhodecode/__init__.py index 0c62e827..aa636536 100644 --- a/rhodecode/__init__.py +++ b/rhodecode/__init__.py @@ -100,7 +100,7 @@ PYRAMID_SETTINGS = {} EXTENSIONS = {} __version__ = ".".join((str(each) for each in VERSION[:3])) -__dbversion__ = 115 # defines current db version for migrations +__dbversion__ = 116 # defines current db version for migrations __license__ = "AGPLv3, and Commercial License" __author__ = "RhodeCode GmbH" __url__ = "https://code.rhodecode.com" diff --git a/rhodecode/apps/repository/tests/test_repo_pullrequests.py b/rhodecode/apps/repository/tests/test_repo_pullrequests.py index 7b9eca13..65eff280 100644 --- a/rhodecode/apps/repository/tests/test_repo_pullrequests.py +++ b/rhodecode/apps/repository/tests/test_repo_pullrequests.py @@ -20,8 +20,8 @@ import os import mock import pytest +from mock.mock import patch, MagicMock -import rhodecode from rhodecode.lib import helpers as h from rhodecode.lib.vcs.backends.base import MergeResponse, MergeFailureReason, Reference from rhodecode.lib.vcs.nodes import FileNode @@ -37,6 +37,7 @@ from rhodecode.model.db import ( ) from rhodecode.model.meta import Session from rhodecode.model.pull_request import PullRequestModel +from rhodecode.model.settings import VcsSettingsModel from rhodecode.model.user import UserModel from rhodecode.model.comment import CommentsModel from rhodecode.tests import ( @@ -44,7 +45,7 @@ from rhodecode.tests import ( TEST_USER_ADMIN_LOGIN, TEST_USER_REGULAR_LOGIN, ) -from rhodecode.tests.fixtures.fixture_utils import PRTestUtility +from rhodecode.tests.fixtures.fixture_utils import PRTestUtility, Backend, temporary_settings from rhodecode.tests.routes import route_path @@ -141,43 +142,19 @@ class TestPullrequestsView(object): source_ref = "branch:{branch}:{commit_id}".format( branch=backend.default_branch_name, commit_id=commit_ids["commit-1"] ) - - response = self.app.post( - route_path("pullrequest_create", repo_name=source.repo_name), - [ - ("source_repo", source_repo_name), - ("source_ref", source_ref), - ("target_repo", target_repo_name), - ("target_ref", target_ref), - ("common_ancestor", commit_ids["initial-commit"]), - ("pullrequest_title", "Title"), - ("pullrequest_desc", "Description"), - ("description_renderer", "markdown"), - ("__start__", "review_members:sequence"), - ("__start__", "reviewer:mapping"), - ("user_id", "1"), - ("__start__", "reasons:sequence"), - ("reason", "Some reason"), - ("__end__", "reasons:sequence"), - ("__start__", "rules:sequence"), - ("__end__", "rules:sequence"), - ("mandatory", "False"), - ("__end__", "reviewer:mapping"), - ("__end__", "review_members:sequence"), - ("__start__", "revisions:sequence"), - ("revisions", commit_ids["commit-1"]), - ("__end__", "revisions:sequence"), - ("user", ""), - ("csrf_token", csrf_token), - ], - status=302, + pr_post_for_params = self._get_pr_create_post_form_params( + csrf_token, source, target, revisions=[commit_ids["commit-1"]] ) + pr_post_for_params.extend( + [ + ("target_ref", target_ref), + ("source_ref", source_ref), + ("common_ancestor", commit_ids["initial-commit"]), + ] + ) + pull_request_id, _ = self._create_pr_and_initial_validation(pr_post_for_params, source) - location = response.headers["Location"] - - pull_request_id = location.rsplit("/", 1)[1] - assert pull_request_id != "new" - pull_request = PullRequest.get(int(pull_request_id)) + pull_request = PullRequest.get(pull_request_id) pull_request_id = pull_request.pull_request_id @@ -354,6 +331,17 @@ class TestPullrequestsView(object): response.mustcontain(cb_context("LINE3")) response.mustcontain(cb_line("LINE4")) + def _create_pr_and_initial_validation(self, pr_post_for_params, source): + response = self.app.post( + route_path("pullrequest_create", repo_name=source.repo_name), + pr_post_for_params, + status=302, + ) + location = response.headers["Location"] + pull_request_id = location.rsplit("/", 1)[1] + assert pull_request_id != "new" + return int(pull_request_id), response + def test_close_status_visibility(self, pr_util, user_util, csrf_token): # Logout response = self.app.post(h.route_path("logout"), params={"csrf_token": csrf_token}) @@ -867,7 +855,109 @@ class TestPullrequestsView(object): ) assert response.status_int == 403 - def test_create_pull_request(self, backend, csrf_token): + def test_new_pull_request_view_close_branch_before_merging_setting_present_for_ee(self, backend, baseapp): + with temporary_settings(baseapp, **{"rhodecode.edition_id": "EE"}): + response = self.app.get(route_path("pullrequest_new", repo_name=backend.repo_name), status=200) + + assert_response = response.assert_response() + assert_response.one_element_exists("#close_branch_before_merging") + + @pytest.mark.parametrize( + "global_close_branch_before_merging_value", + [ + False, + True, + ], + ) + def test_new_pull_request_view_close_branch_before_merging_setting_inherit_global_settings( + self, backend, baseapp, global_close_branch_before_merging_value + ): + def get_repo_settings_inherited(settings_key, default): + if settings_key == f"rhodecode_{backend.repo.repo_type}_close_branch_before_merging": + return global_close_branch_before_merging_value + return None + + with temporary_settings(baseapp, **{"rhodecode.edition_id": "EE"}): + with patch.object(VcsSettingsModel, "get_repo_settings_inherited") as repo_settings_inherited: + settings_mock = MagicMock() + settings_mock.get.side_effect = get_repo_settings_inherited + repo_settings_inherited.return_value = settings_mock + + response = self.app.get(route_path("pullrequest_new", repo_name=backend.repo_name), status=200) + + assert_response = response.assert_response() + if global_close_branch_before_merging_value: + assert_response.element_contain_attribute("#close_branch_before_merging", "checked", "checked") + else: + assert_response.element_contain_no_attribute("#close_branch_before_merging", "checked") + + def test_new_pull_request_view_close_branch_before_merging_setting_not_present_for_ce(self, backend): + response = self.app.get(route_path("pullrequest_new", repo_name=backend.repo_name), status=200) + + assert_response = response.assert_response() + assert_response.no_element_exists("#close_branch_before_merging") + + @pytest.mark.parametrize( + "close_branch_before_merging", + [ + True, + False, + ], + ) + def test_update_pull_request_view_close_branch_before_merging_setting( + self, backend, pr_util, csrf_token, close_branch_before_merging + ): + pull_request = pr_util.create_pull_request(mergeable=True, enable_notifications=False) + pr_id = pull_request.pull_request_id + + data = [ + ("close_branch_before_merging", close_branch_before_merging), + ("csrf_token", csrf_token), + ] + + self.app.post( + route_path( + "pullrequest_update", + repo_name=pull_request.target_repo.scm_instance().name, + pull_request_id=pr_id, + ), + data, + ) + + db_pr = PullRequest.get(pr_id) + + assert db_pr.settings["close_branch_before_merging"] == close_branch_before_merging + + def test_show_pull_request_view_close_branch_before_merging_setting_present_for_ee(self, backend, pr_util, baseapp): + with temporary_settings(baseapp, **{"rhodecode.edition_id": "EE"}): + pull_request = pr_util.create_pull_request(mergeable=True, enable_notifications=False) + + response = self.app.get( + route_path( + "pullrequest_show", + repo_name=pull_request.target_repo.scm_instance().name, + pull_request_id=pull_request.pull_request_id, + ) + ) + + assert_response = response.assert_response() + assert_response.one_element_exists("#close_branch_before_merging") + + def test_show_pull_request_view_close_branch_before_merging_setting_not_present_for_ce(self, backend, pr_util): + pull_request = pr_util.create_pull_request(mergeable=True, enable_notifications=False) + + response = self.app.get( + route_path( + "pullrequest_show", + repo_name=pull_request.target_repo.scm_instance().name, + pull_request_id=pull_request.pull_request_id, + ) + ) + + assert_response = response.assert_response() + assert_response.no_element_exists("#close_branch_before_merging") + + def test_create_pull_request_user_set_test_create_pull_request_to_true(self, backend, csrf_token): commits = [ {"message": "ancestor"}, {"message": "change"}, @@ -877,48 +967,82 @@ class TestPullrequestsView(object): target = backend.create_repo(heads=["ancestor"]) source = backend.create_repo(heads=["change2"]) - response = self.app.post( - route_path("pullrequest_create", repo_name=source.repo_name), + pr_post_for_params = self._get_pr_create_post_form_params( + csrf_token, source, target, revisions=[commit_ids["change"], commit_ids["change2"]] + ) + pr_post_for_params.extend( [ - ("source_repo", source.repo_name), + ("close_branch_before_merging", "true"), ("source_ref", "branch:default:" + commit_ids["change2"]), - ("target_repo", target.repo_name), ("target_ref", "branch:default:" + commit_ids["ancestor"]), ("common_ancestor", commit_ids["ancestor"]), - ("pullrequest_title", "Title"), - ("pullrequest_desc", "Description"), - ("description_renderer", "markdown"), - ("__start__", "review_members:sequence"), - ("__start__", "reviewer:mapping"), - ("user_id", "1"), - ("__start__", "reasons:sequence"), - ("reason", "Some reason"), - ("__end__", "reasons:sequence"), - ("__start__", "rules:sequence"), - ("__end__", "rules:sequence"), - ("mandatory", "False"), - ("__end__", "reviewer:mapping"), - ("__end__", "review_members:sequence"), - ("__start__", "revisions:sequence"), - ("revisions", commit_ids["change"]), - ("revisions", commit_ids["change2"]), - ("__end__", "revisions:sequence"), - ("user", ""), - ("csrf_token", csrf_token), - ], - status=302, + ] + ) + pull_request_id, _ = self._create_pr_and_initial_validation(pr_post_for_params, source) + + pull_request = PullRequest.get(pull_request_id) + + assert len(pull_request.settings) == 1 + assert pull_request.settings["close_branch_before_merging"] + + def _get_pr_create_post_form_params( + self, csrf_token: str, source: Backend, target: Backend, revisions: list, user_id: str = "1" + ): + return [ + ("source_repo", source.repo_name), + ("target_repo", target.repo_name), + ("pullrequest_title", "Title"), + ("pullrequest_desc", "Description"), + ("description_renderer", "markdown"), + ("__start__", "review_members:sequence"), + ("__start__", "reviewer:mapping"), + ("user_id", user_id), + ("__start__", "reasons:sequence"), + ("reason", "Some reason"), + ("__end__", "reasons:sequence"), + ("__start__", "rules:sequence"), + ("__end__", "rules:sequence"), + ("mandatory", "False"), + ("__end__", "reviewer:mapping"), + ("__end__", "review_members:sequence"), + ("__start__", "revisions:sequence"), + *[("revisions", r) for r in revisions], + ("__end__", "revisions:sequence"), + ("user", ""), + ("csrf_token", csrf_token), + ] + + def test_create_pull_request(self, backend, csrf_token): + commits = [ + {"message": "ancestor"}, + {"message": "change"}, + {"message": "change2"}, + ] + commit_ids = backend.create_master_repo(commits) + target = backend.create_repo(heads=["ancestor"]) + source = backend.create_repo(heads=["change2"]) + pr_post_for_params = self._get_pr_create_post_form_params( + csrf_token, source, target, revisions=[commit_ids["change"], commit_ids["change2"]] + ) + pr_post_for_params.extend( + [ + ("source_ref", "branch:default:" + commit_ids["change2"]), + ("target_ref", "branch:default:" + commit_ids["ancestor"]), + ("common_ancestor", commit_ids["ancestor"]), + ] ) - location = response.headers["Location"] - pull_request_id = location.rsplit("/", 1)[1] - assert pull_request_id != "new" - pull_request = PullRequest.get(int(pull_request_id)) + pull_request_id, _ = self._create_pr_and_initial_validation(pr_post_for_params, source) + + pull_request = PullRequest.get(pull_request_id) # check that we have now both revisions assert pull_request.revisions == [commit_ids["change2"], commit_ids["change"]] assert pull_request.source_ref == "branch:default:" + commit_ids["change2"] expected_target_ref = "branch:default:" + commit_ids["ancestor"] assert pull_request.target_ref == expected_target_ref + assert len(pull_request.settings) == 1 + assert not pull_request.settings["close_branch_before_merging"] def test_reviewer_notifications(self, backend, csrf_token): # We have to use the app.post for this test, so it will create the @@ -945,42 +1069,20 @@ class TestPullrequestsView(object): target = backend.create_repo(heads=["ancestor-child"]) source = backend.create_repo(heads=["change"]) - response = self.app.post( - route_path("pullrequest_create", repo_name=source.repo_name), + pr_post_for_params = self._get_pr_create_post_form_params( + csrf_token, source, target, user_id="2", revisions=[commit_ids["change"]] + ) + pr_post_for_params.extend( [ - ("source_repo", source.repo_name), ("source_ref", "branch:default:" + commit_ids["change"]), - ("target_repo", target.repo_name), ("target_ref", "branch:default:" + commit_ids["ancestor-child"]), ("common_ancestor", commit_ids["ancestor"]), - ("pullrequest_title", "Title"), - ("pullrequest_desc", "Description"), - ("description_renderer", "markdown"), - ("__start__", "review_members:sequence"), - ("__start__", "reviewer:mapping"), - ("user_id", "2"), - ("__start__", "reasons:sequence"), - ("reason", "Some reason"), - ("__end__", "reasons:sequence"), - ("__start__", "rules:sequence"), - ("__end__", "rules:sequence"), - ("mandatory", "False"), - ("__end__", "reviewer:mapping"), - ("__end__", "review_members:sequence"), - ("__start__", "revisions:sequence"), - ("revisions", commit_ids["change"]), - ("__end__", "revisions:sequence"), - ("user", ""), - ("csrf_token", csrf_token), - ], - status=302, + ] ) - location = response.headers["Location"] + pull_request_id, _ = self._create_pr_and_initial_validation(pr_post_for_params, source) - pull_request_id = location.rsplit("/", 1)[1] - assert pull_request_id != "new" - pull_request = PullRequest.get(int(pull_request_id)) + pull_request = PullRequest.get(pull_request_id) # Check that a notification was made notifications = Notification.query().filter( @@ -1023,43 +1125,20 @@ class TestPullrequestsView(object): commit_ids = backend.create_master_repo(commits) target = backend.create_repo(heads=["ancestor-child"]) source = backend.create_repo(heads=["change"]) - - response = self.app.post( - route_path("pullrequest_create", repo_name=source.repo_name), + pr_post_for_params = self._get_pr_create_post_form_params( + csrf_token, source, target, revisions=[commit_ids["change"]] + ) + pr_post_for_params.extend( [ - ("source_repo", source.repo_name), ("source_ref", "branch:default:" + commit_ids["change"]), - ("target_repo", target.repo_name), ("target_ref", "branch:default:" + commit_ids["ancestor-child"]), ("common_ancestor", commit_ids["ancestor"]), - ("pullrequest_title", "Title"), - ("pullrequest_desc", "Description"), - ("description_renderer", "markdown"), - ("__start__", "review_members:sequence"), - ("__start__", "reviewer:mapping"), - ("user_id", "1"), - ("__start__", "reasons:sequence"), - ("reason", "Some reason"), - ("__end__", "reasons:sequence"), - ("__start__", "rules:sequence"), - ("__end__", "rules:sequence"), - ("mandatory", "False"), - ("__end__", "reviewer:mapping"), - ("__end__", "review_members:sequence"), - ("__start__", "revisions:sequence"), - ("revisions", commit_ids["change"]), - ("__end__", "revisions:sequence"), - ("user", ""), - ("csrf_token", csrf_token), - ], - status=302, + ] ) - location = response.headers["Location"] + pull_request_id, response = self._create_pr_and_initial_validation(pr_post_for_params, source) - pull_request_id = location.rsplit("/", 1)[1] - assert pull_request_id != "new" - pull_request = PullRequest.get(int(pull_request_id)) + pull_request = PullRequest.get(pull_request_id) # target_ref has to point to the ancestor's commit_id in order to # show the correct diff @@ -1153,7 +1232,7 @@ class TestPullrequestsView(object): mods = [ ( "_pre_push_hook", - f""" + """ return HookResponse(1, 'HOOK_TEST_FORBIDDEN') """, ) diff --git a/rhodecode/apps/repository/views/repo_pull_requests.py b/rhodecode/apps/repository/views/repo_pull_requests.py index 9b24e585..69edc3d9 100644 --- a/rhodecode/apps/repository/views/repo_pull_requests.py +++ b/rhodecode/apps/repository/views/repo_pull_requests.py @@ -166,6 +166,7 @@ class RepoPullRequestsView(RepoAppView, DataGridAppView): "comments": _render("pullrequest_comments", comments_count), "comments_raw": comments_count, "closed": pr.is_closed(), + "settings": pr.settings, } ) @@ -934,6 +935,9 @@ class RepoPullRequestsView(RepoAppView, DataGridAppView): } c.default_source_ref = selected_source_ref + close_branch_before_merging_key = "rhodecode_%s_close_branch_before_merging" % source_repo.repo_type + c.repo_close_branch_before_merging = self._get_repo_setting(source_repo, close_branch_before_merging_key) + return self._get_template_context(c) @LoginRequired() @@ -1220,6 +1224,7 @@ class RepoPullRequestsView(RepoAppView, DataGridAppView): description = _form["pullrequest_desc"] description_renderer = _form["description_renderer"] + settings = {"close_branch_before_merging": _form["close_branch_before_merging"]} try: pull_request = PullRequestModel().create( @@ -1237,6 +1242,7 @@ class RepoPullRequestsView(RepoAppView, DataGridAppView): description_renderer=description_renderer, reviewer_data=reviewer_rules, auth_user=self._rhodecode_user, + settings=settings, ) Session().commit() @@ -1283,6 +1289,7 @@ class RepoPullRequestsView(RepoAppView, DataGridAppView): controls = peppercorn.parse(self.request.POST.items()) force_refresh = str2bool(self.request.POST.get("force_refresh", "false")) do_update_commits = str2bool(self.request.POST.get("update_commits", "false")) + do_update_branch_close = "close_branch_before_merging" in self.request.POST if "review_members" in controls: self._update_reviewers( @@ -1321,6 +1328,8 @@ class RepoPullRequestsView(RepoAppView, DataGridAppView): ) elif str2bool(self.request.POST.get("edit_pull_request", "false")): self._edit_pull_request(pull_request) + elif do_update_branch_close: + self._update_settings(pull_request) else: log.error("Unhandled update data.") raise HTTPBadRequest() @@ -1328,6 +1337,14 @@ class RepoPullRequestsView(RepoAppView, DataGridAppView): return {"response": True, "redirect_url": redirect_url} raise HTTPForbidden() + def _update_settings(self, pull_request): + try: + close_branch_before_merging = str2bool(self.request.POST.get("close_branch_before_merging", "false")) + PullRequestModel().update_settings(pull_request, close_branch_before_merging) + except ValueError: + msg = self.request.translate("Cannot update closed pull requests.") + h.flash(msg, category="error") + def _edit_pull_request(self, pull_request): """ Edit title and description diff --git a/rhodecode/lib/dbmigrate/versions/116_version_5_7_0.py b/rhodecode/lib/dbmigrate/versions/116_version_5_7_0.py new file mode 100644 index 00000000..4ab47a27 --- /dev/null +++ b/rhodecode/lib/dbmigrate/versions/116_version_5_7_0.py @@ -0,0 +1,75 @@ +import json +import logging +from sqlalchemy import * +from sqlalchemy.engine import reflection + +from alembic.migration import MigrationContext +from alembic.operations import Operations + +from rhodecode.lib.dbmigrate.versions import _reset_base +from rhodecode.lib.jsonalchemy import MutationObj, JsonType +from rhodecode.model import meta, init_model_encryption + + +def upgrade(migrate_engine): + """ + Upgrade operations go here. + Don't create your own engine; bind migrate_engine to your metadata + """ + _reset_base(migrate_engine) + + from rhodecode.lib.dbmigrate.schema import db_4_20_0_0 as db + + init_model_encryption(db) + + context = MigrationContext.configure(migrate_engine.connect()) + op = Operations(context) + + pr_table = db.PullRequest.__table__ + with op.batch_alter_table(pr_table.name) as batch_op: + new_column = Column( + "settings_json", + MutationObj.as_mutable( + JsonType(dialect_map=dict(mysql=UnicodeText(16384))), + ), + default=dict, + ) + batch_op.add_column(new_column) + + pr_version_table = db.PullRequestVersion.__table__ + with op.batch_alter_table(pr_version_table.name) as batch_op: + new_column = Column( + "settings_json", + MutationObj.as_mutable( + JsonType(dialect_map=dict(mysql=UnicodeText(16384))), + ), + default=dict, + ) + batch_op.add_column(new_column) + + _inherit_settings(db, meta.Session, op) + + +def downgrade(migrate_engine): + pass + + +def _inherit_settings(models, _SESSION, op): + for pr in _SESSION.query(models.PullRequest).all(): + repo = pr.target_repo + repo_type = repo.repo_type + close_branch_before_merging = False + + if repo_type in ["git", "hg"]: + close_branch_before_merging = getattr(repo, f"{repo_type}_close_branch_before_merging", False) + + json_settings = {"close_branch_before_merging": close_branch_before_merging} + params = {"id": pr.pull_request_id, "value": json.dumps(json_settings)} + query = text( + """UPDATE pull_requests + SET settings_json = :value + WHERE pull_request_id = :id""" + ).bindparams(**params) + op.execute(query) + + _SESSION().commit() diff --git a/rhodecode/model/db.py b/rhodecode/model/db.py index 1227902d..8c781a53 100644 --- a/rhodecode/model/db.py +++ b/rhodecode/model/db.py @@ -77,7 +77,6 @@ from zope.cachedescriptors.property import Lazy as LazyProperty from webhelpers2.text import remove_formatting -from rhodecode import ConfigGet from rhodecode.lib.str_utils import safe_bytes from rhodecode.translation import _ from rhodecode.lib.vcs import get_vcs_instance, VCSError @@ -108,6 +107,8 @@ from rhodecode.lib.exceptions import ArtifactMetadataDuplicate, ArtifactMetadata from rhodecode.lib.pyramid_utils import get_current_request from rhodecode.model.meta import Base, Session +DEFAULT_JSON_OBJ_SIZE = 16384 + URL_SEP = "/" log = logging.getLogger(__name__) @@ -4443,13 +4444,23 @@ class _PullRequestBase(BaseModel): _last_merge_target_rev = Column("last_merge_other_rev", String(40), nullable=True) _last_merge_status = Column("merge_status", Integer(), nullable=True) last_merge_metadata = Column( - "last_merge_metadata", MutationObj.as_mutable(JsonType(dialect_map=dict(mysql=UnicodeText(16384)))) + "last_merge_metadata", + MutationObj.as_mutable(JsonType(dialect_map=dict(mysql=UnicodeText(DEFAULT_JSON_OBJ_SIZE)))), ) merge_rev = Column("merge_rev", String(40), nullable=True) reviewer_data = Column( - "reviewer_data_json", MutationObj.as_mutable(JsonType(dialect_map=dict(mysql=UnicodeText(16384)))) + "reviewer_data_json", + MutationObj.as_mutable(JsonType(dialect_map=dict(mysql=UnicodeText(DEFAULT_JSON_OBJ_SIZE)))), + ) + + settings = Column( + "settings_json", + MutationObj.as_mutable( + JsonType(dialect_map=dict(mysql=UnicodeText(DEFAULT_JSON_OBJ_SIZE))), + ), + default=dict, ) @property @@ -4711,6 +4722,7 @@ class PullRequest(Base, _PullRequestBase): attrs.target_ref_parts = pull_request_obj.target_ref_parts attrs.revisions = pull_request_obj.revisions attrs.common_ancestor_id = pull_request_obj.common_ancestor_id + attrs.settings = pull_request_obj.settings attrs.shadow_merge_ref = org_pull_request_obj.shadow_merge_ref attrs.reviewer_data = org_pull_request_obj.reviewer_data attrs.reviewer_data_json = org_pull_request_obj.reviewer_data_json @@ -4798,6 +4810,26 @@ class PullRequest(Base, _PullRequestBase): return self.versions_count +@event.listens_for(PullRequest, "before_insert") +def _init_pr_default_settings(mapper, connection, pull_request): + if not pull_request.settings: + pull_request.settings = { + "close_branch_before_merging": _inherit_global_settings(pull_request), + } + + +def _inherit_global_settings(pull_request): + from rhodecode.model.settings import VcsSettingsModel # handle circular dependency issue + + repo_type = pull_request.target_repo.repo_type + settings_model = VcsSettingsModel(repo=pull_request.target_repo) + settings = settings_model.get_general_settings() + key = "rhodecode_{}_close_branch_before_merging" + if repo_type in ["hg", "git"]: + return settings.get(key.format(repo_type), False) + return False + + class PullRequestVersion(Base, _PullRequestBase): __tablename__ = "pull_request_versions" __table_args__ = (base_table_args,) @@ -4860,7 +4892,9 @@ class PullRequestReviewers(Base, BaseModel): pull_requests_reviewers_id = Column("pull_requests_reviewers_id", Integer(), nullable=False, primary_key=True) pull_request_id = Column("pull_request_id", Integer(), ForeignKey("pull_requests.pull_request_id"), nullable=False) user_id = Column("user_id", Integer(), ForeignKey("users.user_id"), nullable=True) - _reasons = Column("reason", MutationList.as_mutable(JsonType("list", dialect_map=dict(mysql=UnicodeText(16384))))) + _reasons = Column( + "reason", MutationList.as_mutable(JsonType("list", dialect_map=dict(mysql=UnicodeText(DEFAULT_JSON_OBJ_SIZE)))) + ) mandatory = Column("mandatory", Boolean(), nullable=False, default=False) role = Column("role", Unicode(255), nullable=True, default=ROLE_REVIEWER) @@ -4868,7 +4902,7 @@ class PullRequestReviewers(Base, BaseModel): user = relationship("User") pull_request = relationship("PullRequest", back_populates="reviewers") - rule_data = Column("rule_data_json", JsonType(dialect_map=dict(mysql=UnicodeText(16384)))) + rule_data = Column("rule_data_json", JsonType(dialect_map=dict(mysql=UnicodeText(DEFAULT_JSON_OBJ_SIZE)))) def rule_user_group_data(self): """ @@ -5233,7 +5267,9 @@ class Integration(Base, BaseModel): name = Column("name", String(255), nullable=False) child_repos_only = Column("child_repos_only", Boolean(), nullable=False, default=False) - settings = Column("settings_json", MutationObj.as_mutable(JsonType(dialect_map=dict(mysql=UnicodeText(16384))))) + settings = Column( + "settings_json", MutationObj.as_mutable(JsonType(dialect_map=dict(mysql=UnicodeText(DEFAULT_JSON_OBJ_SIZE)))) + ) repo_id = Column("repo_id", Integer(), ForeignKey("repositories.repo_id"), nullable=True, unique=None, default=None) repo = relationship("Repository", lazy="joined", back_populates="integrations") diff --git a/rhodecode/model/forms.py b/rhodecode/model/forms.py index 695886b2..5fd7a62f 100644 --- a/rhodecode/model/forms.py +++ b/rhodecode/model/forms.py @@ -647,6 +647,7 @@ def PullRequestForm(localizer, repo_id): pullrequest_title = v.UnicodeString(strip=True, required=True, min=1, max=255) pullrequest_desc = v.UnicodeString(strip=True, required=False) description_renderer = v.UnicodeString(strip=True, required=False) + close_branch_before_merging = v.StringBoolean(if_missing=False) return _PullRequestForm diff --git a/rhodecode/model/pull_request.py b/rhodecode/model/pull_request.py index dedfb409..71156d36 100644 --- a/rhodecode/model/pull_request.py +++ b/rhodecode/model/pull_request.py @@ -31,6 +31,10 @@ import urllib.error import collections import dataclasses as dataclasses +from copy import deepcopy + +from pyramid.threadlocal import get_current_registry + from rhodecode.lib.pyramid_utils import get_current_request from rhodecode.lib.vcs.nodes import FileNode @@ -809,6 +813,7 @@ class PullRequestModel(BaseModel): reviewer_data=None, translator=None, auth_user=None, + settings=None, ): translator = translator or get_current_request().translate @@ -830,6 +835,8 @@ class PullRequestModel(BaseModel): pull_request.reviewer_data = reviewer_data pull_request.pull_request_state = pull_request.STATE_CREATING pull_request.common_ancestor_id = common_ancestor_id + if self._settings_valid(settings): + pull_request.settings = settings Session().add(pull_request) Session().flush() @@ -938,6 +945,19 @@ class PullRequestModel(BaseModel): return pull_request + def _settings_valid(self, settings): + if not settings: + return False + if not isinstance(settings, dict): + return False + if len(settings) > 1: + return False + if "close_branch_before_merging" not in settings: + return False + if not isinstance(settings["close_branch_before_merging"], bool): + return False + return True + def trigger_pull_request_hook(self, pull_request, user, action, data=None): pull_request = self.__get_pull_request(pull_request) target_scm = pull_request.target_repo.scm_instance() @@ -1470,6 +1490,19 @@ class PullRequestModel(BaseModel): renderer = RstTemplateRenderer() return renderer.render("pull_request_update.mako", **params) + def update_settings(self, pull_request: PullRequest, close_branch_before_merging: bool): + pull_request = self.__get_pull_request(pull_request) + if pull_request.is_closed(): + raise ValueError("This pull request is closed") + + if pull_request.settings["close_branch_before_merging"] == close_branch_before_merging: + return + + settings = deepcopy(pull_request.settings) # need to copy, otherwise SQLalchemy not tracking changes + settings["close_branch_before_merging"] = close_branch_before_merging + pull_request.settings = settings + Session().commit() + def edit(self, pull_request, title, description, description_renderer, user): pull_request = self.__get_pull_request(pull_request) old_data = pull_request.get_api_data(with_merge_state=False) @@ -2191,14 +2224,20 @@ class PullRequestModel(BaseModel): user_name = getattr(user, user_name_attr) return user_name - def _close_branch_before_merging(self, pull_request): + def _close_branch_before_merging(self, pull_request: PullRequest): repo_type = pull_request.target_repo.repo_type - if repo_type == "hg": - return self._get_general_setting(pull_request, "rhodecode_hg_close_branch_before_merging") - elif repo_type == "git": - return self._get_general_setting(pull_request, "rhodecode_git_close_branch_before_merging") + if repo_type not in ["hg", "git"]: + return False - return False + registry = get_current_registry() + is_enterprise = registry.settings.get("rhodecode.edition_id") == "EE" + + if is_enterprise and pull_request.settings and "close_branch_before_merging" in pull_request.settings: + # this feature is available only for the EE edition + return pull_request.settings["close_branch_before_merging"] + + key = "rhodecode_{}_close_branch_before_merging".format(repo_type) + return self._get_general_setting(pull_request, key) def _get_general_setting(self, pull_request, settings_key, default=False): settings_model = VcsSettingsModel(repo=pull_request.target_repo) diff --git a/rhodecode/public/css/main.less b/rhodecode/public/css/main.less index fe1b3515..fab77fea 100644 --- a/rhodecode/public/css/main.less +++ b/rhodecode/public/css/main.less @@ -1976,6 +1976,10 @@ BIN_FILENODE = 7 } } +.pull-request-settings { + margin: 2px 7px; +} + .pull-request-merge ul { padding: 0px 0px; } diff --git a/rhodecode/public/js/src/rhodecode/pullrequests.js b/rhodecode/public/js/src/rhodecode/pullrequests.js index 550be5ac..feceb658 100644 --- a/rhodecode/public/js/src/rhodecode/pullrequests.js +++ b/rhodecode/public/js/src/rhodecode/pullrequests.js @@ -528,6 +528,13 @@ var autoCompleteHandler = function (inputId, controller, role) { } } +let updateCloseBranchSetting = function(repo_name, pull_request_id, close_branch_before_merging) { + const postData = { + 'close_branch_before_merging': close_branch_before_merging, + }; + _updatePullRequest(repo_name, pull_request_id, postData); +} + /** * Reviewer autocomplete */ diff --git a/rhodecode/templates/base/vcs_settings.mako b/rhodecode/templates/base/vcs_settings.mako index 6bd45548..003f1317 100644 --- a/rhodecode/templates/base/vcs_settings.mako +++ b/rhodecode/templates/base/vcs_settings.mako @@ -296,7 +296,7 @@
${h.checkbox('rhodecode_git_close_branch_before_merging' + suffix, 'True', **kwargs)} - +
${_('Delete branch after merging it into destination branch.')} diff --git a/rhodecode/templates/pullrequests/pullrequest.mako b/rhodecode/templates/pullrequests/pullrequest.mako index 0942434e..4d2372e3 100644 --- a/rhodecode/templates/pullrequests/pullrequest.mako +++ b/rhodecode/templates/pullrequests/pullrequest.mako @@ -216,6 +216,19 @@
+ % if c.rhodecode_edition_id == 'EE': +
+ ${h.checkbox('close_branch_before_merging', + checked=c.repo_close_branch_before_merging, value=True)} + +
+ % endif diff --git a/rhodecode/templates/pullrequests/pullrequest_merge_checks.mako b/rhodecode/templates/pullrequests/pullrequest_merge_checks.mako index 8386d7eb..c1712b04 100644 --- a/rhodecode/templates/pullrequests/pullrequest_merge_checks.mako +++ b/rhodecode/templates/pullrequests/pullrequest_merge_checks.mako @@ -66,6 +66,21 @@ ${h.end_form()} + % if c.rhodecode_edition_id == 'EE': +
+ ${h.checkbox('close_branch_before_merging', checked=c.pull_request.settings.get("close_branch_before_merging", False))} + +
+ % endif +
${_('refresh checks')}
@@ -80,3 +95,14 @@ + + diff --git a/rhodecode/tests/fixtures/fixture_utils.py b/rhodecode/tests/fixtures/fixture_utils.py index 9599e549..ff0ab6d6 100644 --- a/rhodecode/tests/fixtures/fixture_utils.py +++ b/rhodecode/tests/fixtures/fixture_utils.py @@ -26,6 +26,9 @@ import socket import subprocess import time import uuid +from contextlib import contextmanager +from copy import copy, deepcopy + import dateutil.tz import logging import functools @@ -172,6 +175,17 @@ def plain_http_environ(): } +@contextmanager +def temporary_settings(baseapp, **overrides): + original = deepcopy(baseapp.config.registry.settings) + try: + baseapp.config.registry.settings.update(overrides) + yield + finally: + baseapp.config.registry.settings.clear() + baseapp.config.registry.settings.update(original) + + @pytest.fixture(scope="session") def baseapp(request, ini_config, http_environ_session, available_port_factory, vcsserver_factory, celery_factory): from rhodecode.lib.config_utils import get_app_config diff --git a/rhodecode/tests/utils.py b/rhodecode/tests/utils.py index 684d1e9d..6ad0fb33 100644 --- a/rhodecode/tests/utils.py +++ b/rhodecode/tests/utils.py @@ -284,6 +284,14 @@ class AssertResponse(object): element = self.get_element(css_selector) assert expected_content in element.text_content() + def element_contain_attribute(self, css_selector, attr_name, attr_value): + element = self.get_element(css_selector) + assert element.attrib.get(attr_name) == attr_value + + def element_contain_no_attribute(self, css_selector, attr_name): + element = self.get_element(css_selector) + assert not element.attrib.get(attr_name) + def element_value_contains(self, css_selector, expected_content): element = self.get_element(css_selector) assert expected_content in element.value