Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
127 changes: 70 additions & 57 deletions comfy_cli/git_utils.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,3 @@
import os
import subprocess

from rich.console import Console
Expand All @@ -24,6 +23,40 @@ def sanitize_for_local_branch(branch_name: str) -> str:
return sanitized or "unknown"


def _print_git_error(
title: str,
panel_title: str,
context: str,
details: list[str],
exc: subprocess.CalledProcessError,
) -> None:
"""Render a git ``CalledProcessError`` as a rich error panel.

Shared by ``git_checkout_tag`` and ``checkout_pr`` so the two failure
displays stay in lock-step. ``context`` is the bold headline line, each
``details`` entry is an italic follow-up line, and the command's stderr
(when present) is appended verbatim.
"""
error_message = Text()
error_message.append(title, style="bold red on white")
error_message.append(f"\n\n{context}", style="bold yellow")
for detail in details:
error_message.append(f"\n{detail}", style="italic")

if exc.stderr:
error_message.append("\n\nError output:", style="bold red")
error_message.append(f"\n{exc.stderr}", style="italic yellow")

console.print(
Panel(
error_message,
title=f"[bold white on red]{panel_title}[/bold white on red]",
border_style="red",
expand=False,
)
)


def git_checkout_tag(repo_path: str, tag: str) -> bool:
"""
Checkout a specific Git tag in the given repository.
Expand All @@ -39,79 +72,64 @@ def git_checkout_tag(repo_path: str, tag: str) -> bool:
:param tag: The tag to checkout
:return: True if the checkout succeeds, False if any git command failed.
"""
original_dir = os.getcwd()
try:
# Change to the repository directory

os.chdir(repo_path)

# Skip the network fetch when the tag is already present locally.
tag_present_locally = (
subprocess.run(
["git", "rev-parse", "--verify", f"refs/tags/{tag}"],
cwd=repo_path,
capture_output=True,
text=True,
check=False,
).returncode
== 0
)
if not tag_present_locally:
subprocess.run(["git", "fetch", "--tags"], check=True, capture_output=True, text=True)
subprocess.run(["git", "fetch", "--tags"], cwd=repo_path, check=True, capture_output=True, text=True)

# Checkout the specified tag
subprocess.run(["git", "checkout", tag], check=True, capture_output=True, text=True)
subprocess.run(["git", "checkout", tag], cwd=repo_path, check=True, capture_output=True, text=True)

console.print(f"[bold green]Successfully checked out tag: [cyan]{tag}[/cyan][/bold green]")

return True
except subprocess.CalledProcessError as e:
error_message = Text()
error_message.append("Git Checkout Error", style="bold red on white")
error_message.append("\n\nFailed to checkout tag: ", style="bold yellow")
error_message.append(f"[cyan]{tag}[/cyan]")
error_message.append("\n\nError details:", style="bold red")
error_message.append(f"\n{str(e)}", style="italic")

if e.stderr:
error_message.append("\n\nError output:", style="bold red")
error_message.append(f"\n{e.stderr}", style="italic yellow")

console.print(
Panel(
error_message,
title="[bold white on red]Git Checkout Failed[/bold white on red]",
border_style="red",
expand=False,
)
_print_git_error(
title="Git Checkout Error",
panel_title="Git Checkout Failed",
context=f"Failed to checkout tag: {tag}",
details=[f"Error details:\n{e}"],
exc=e,
)

return False
finally:
# Ensure we always return to the original directory
os.chdir(original_dir)


def checkout_pr(repo_path: str, pr_info: PRInfo) -> bool:
original_dir = os.getcwd()

try:
os.chdir(repo_path)

if pr_info.is_fork:
remote_name = f"pr-{pr_info.number}-{pr_info.user}"

result = subprocess.run(["git", "remote", "get-url", remote_name], capture_output=True, text=True)
result = subprocess.run(
["git", "remote", "get-url", remote_name], cwd=repo_path, capture_output=True, text=True
)

if result.returncode != 0:
subprocess.run(
["git", "remote", "add", remote_name, pr_info.head_repo_url],
cwd=repo_path,
check=True,
capture_output=True,
text=True,
)

subprocess.run(
["git", "fetch", remote_name, pr_info.head_branch], check=True, capture_output=True, text=True
# ``--`` guards against a fork-controlled branch name beginning with ``-`` being
# parsed as a git option (argument injection); ``comfy install --pr`` runs against forks.
["git", "fetch", remote_name, "--", pr_info.head_branch],
cwd=repo_path,
check=True,
capture_output=True,
text=True,
)

# fix: "feature/add-support" -> "pr-123-feature-add-support"
Expand All @@ -120,19 +138,28 @@ def checkout_pr(repo_path: str, pr_info: PRInfo) -> bool:

subprocess.run(
["git", "checkout", "-B", local_branch, f"{remote_name}/{pr_info.head_branch}"],
cwd=repo_path,
check=True,
capture_output=True,
text=True,
)

else:
subprocess.run(["git", "fetch", "origin", pr_info.head_branch], check=True, capture_output=True, text=True)
subprocess.run(
# ``--`` guards against a branch name beginning with ``-`` being parsed as a git option.
["git", "fetch", "origin", "--", pr_info.head_branch],
cwd=repo_path,
check=True,
capture_output=True,
text=True,
)

sanitized_branch = sanitize_for_local_branch(pr_info.head_branch)
local_branch = f"pr-{pr_info.number}-{sanitized_branch}"

subprocess.run(
["git", "checkout", "-B", local_branch, f"origin/{pr_info.head_branch}"],
cwd=repo_path,
check=True,
capture_output=True,
text=True,
Expand All @@ -143,25 +170,11 @@ def checkout_pr(repo_path: str, pr_info: PRInfo) -> bool:
return True

except subprocess.CalledProcessError as e:
error_message = Text()
error_message.append("Git PR Checkout Error", style="bold red on white")
error_message.append(f"\n\nFailed to checkout PR #{pr_info.number}", style="bold yellow")
error_message.append(f"\nTitle: {pr_info.title}", style="italic")
error_message.append(f"\nBranch: {pr_info.head_branch}", style="italic")

if e.stderr:
error_message.append("\n\nError output:", style="bold red")
error_message.append(f"\n{e.stderr}", style="italic yellow")

console.print(
Panel(
error_message,
title="[bold white on red]PR Checkout Failed[/bold white on red]",
border_style="red",
expand=False,
)
_print_git_error(
title="Git PR Checkout Error",
panel_title="PR Checkout Failed",
context=f"Failed to checkout PR #{pr_info.number}",
details=[f"Title: {pr_info.title}", f"Branch: {pr_info.head_branch}"],
exc=e,
)
return False

finally:
os.chdir(original_dir)
34 changes: 14 additions & 20 deletions tests/comfy_cli/command/github/test_pr.py
Original file line number Diff line number Diff line change
Expand Up @@ -238,12 +238,8 @@ class TestGitOperations:
"""Test Git operations for PR checkout"""

@patch("subprocess.run")
@patch("os.chdir")
@patch("os.getcwd")
def test_checkout_pr_fork_success(self, mock_getcwd, mock_chdir, mock_subprocess, sample_pr_info):
def test_checkout_pr_fork_success(self, mock_subprocess, sample_pr_info):
"""Test successful checkout of PR from fork"""
mock_getcwd.return_value = "/original/dir"

mock_subprocess.side_effect = [
subprocess.CompletedProcess([], 1),
subprocess.CompletedProcess([], 0),
Expand All @@ -261,14 +257,15 @@ def test_checkout_pr_fork_success(self, mock_getcwd, mock_chdir, mock_subprocess
assert "remote" in calls[1][0][0]
assert "fetch" in calls[2][0][0]
assert "checkout" in calls[3][0][0]
# ``--`` separates options from the fork-controlled branch refspec (argument-injection guard).
fetch_argv = calls[2][0][0]
assert fetch_argv[-2:] == ["--", "load-3d-nodes"]
# Every git command runs against repo_path via cwd= (no process chdir).
assert all(call.kwargs.get("cwd") == "/repo/path" for call in calls)

@patch("subprocess.run")
@patch("os.chdir")
@patch("os.getcwd")
def test_checkout_pr_non_fork_success(self, mock_getcwd, mock_chdir, mock_subprocess):
def test_checkout_pr_non_fork_success(self, mock_subprocess):
"""Test successful checkout of PR from same repo"""
mock_getcwd.return_value = "/original/dir"

pr_info = PRInfo(
number=123,
head_repo_url="https://github.com/comfyanonymous/ComfyUI.git",
Expand All @@ -289,14 +286,14 @@ def test_checkout_pr_non_fork_success(self, mock_getcwd, mock_chdir, mock_subpro

assert result is True
assert mock_subprocess.call_count == 2
# ``--`` separates options from the branch refspec (argument-injection guard).
fetch_argv = mock_subprocess.call_args_list[0][0][0]
assert fetch_argv[-2:] == ["--", "feature-branch"]
assert all(call.kwargs.get("cwd") == "/repo/path" for call in mock_subprocess.call_args_list)

@patch("subprocess.run")
@patch("os.chdir")
@patch("os.getcwd")
def test_checkout_pr_git_failure(self, mock_getcwd, mock_chdir, mock_subprocess, sample_pr_info):
def test_checkout_pr_git_failure(self, mock_subprocess, sample_pr_info):
"""Test Git operation failure"""
mock_getcwd.return_value = "/original/dir"

error = subprocess.CalledProcessError(1, "git", stderr="Permission denied")
mock_subprocess.side_effect = error

Expand Down Expand Up @@ -553,12 +550,8 @@ def test_fetch_pr_info_with_github_token(self, mock_get):
assert headers["Authorization"] == "Bearer test-token"

@patch("subprocess.run")
@patch("os.chdir")
@patch("os.getcwd")
def test_checkout_pr_remote_already_exists(self, mock_getcwd, mock_chdir, mock_subprocess, sample_pr_info):
def test_checkout_pr_remote_already_exists(self, mock_subprocess, sample_pr_info):
"""Test checkout when remote already exists"""
mock_getcwd.return_value = "/dir"

mock_subprocess.side_effect = [
subprocess.CompletedProcess([], 0),
subprocess.CompletedProcess([], 0),
Expand All @@ -569,6 +562,7 @@ def test_checkout_pr_remote_already_exists(self, mock_getcwd, mock_chdir, mock_s

assert result is True
assert mock_subprocess.call_count == 3
assert all(call.kwargs.get("cwd") == "/repo" for call in mock_subprocess.call_args_list)


class TestGetLatestRelease:
Expand Down
Loading