Files
coding-agent-gitea/tests/test_pr_tools.py
T
meeks dfdd8c0931 fix: improve prompts, error messages, and workspace concurrency
- Externalize coordinator, notification, and planning prompts to separate files
- Add workspace mutex for concurrent file operations
- Improve error messages across file_tools, issue_tools, and pr_tools
- Add logging to tool modules for better debugging
- Update tests to match new error message strings
2026-08-01 20:01:46 +02:00

263 lines
9.1 KiB
Python

import json
from typing import Any
from unittest.mock import MagicMock
from gitea.client import GiteaClient
from gitea.models import PullRequestModel, CommentModel, RepositoryModel
from gitea.tools.pr_tools import PRTools
def _create_mock_client() -> MagicMock:
"""Create a mock GiteaClient with sub-client attributes."""
mock_client: MagicMock = MagicMock(spec=GiteaClient)
mock_client.prs = MagicMock()
mock_client.repos = MagicMock()
return mock_client
def test_get_pull_request_success() -> None:
mock_client = _create_mock_client()
pr: PullRequestModel = PullRequestModel(number=1, title="Test PR", state="open")
mock_client.prs.get_pull_request.return_value = pr
pr_tools: PRTools = PRTools(mock_client)
res: PullRequestModel = pr_tools.get_pull_request("owner", "repo", 1)
assert isinstance(res, PullRequestModel)
assert res.number == 1
assert res.title == "Test PR"
mock_client.prs.get_pull_request.assert_called_once_with("owner", "repo", 1)
def test_get_pull_request_failure() -> None:
mock_client = _create_mock_client()
mock_client.prs.get_pull_request.side_effect = Exception("API Error")
pr_tools: PRTools = PRTools(mock_client)
try:
pr_tools.get_pull_request("owner", "repo", 1)
assert False, "Expected Exception"
except Exception as e:
assert str(e) == "API Error"
def test_close_pull_request_success() -> None:
mock_client = _create_mock_client()
mock_client.prs.close_pull_request.return_value = PullRequestModel(
number=1, state="closed"
)
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.close_pull_request("owner", "repo", 1)
assert res == "Pull request #1 closed successfully."
mock_client.prs.close_pull_request.assert_called_once_with("owner", "repo", 1)
def test_close_pull_request_failure() -> None:
mock_client = _create_mock_client()
mock_client.prs.close_pull_request.side_effect = Exception("API Error")
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.close_pull_request("owner", "repo", 1)
assert "Could not close PR" in res
assert "owner/repo" in res
def test_get_pull_request_comments_success() -> None:
mock_client = _create_mock_client()
comment: CommentModel = CommentModel(id=123, body="Comment body")
mock_client.prs.get_pull_request_comments.return_value = [comment]
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.get_pull_request_comments("owner", "repo", 1)
data: list[dict[str, Any]] = json.loads(res)
assert len(data) == 1
assert data[0]["body"] == "Comment body"
def test_get_pull_request_comments_failure() -> None:
mock_client = _create_mock_client()
mock_client.prs.get_pull_request_comments.side_effect = Exception("API Error")
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.get_pull_request_comments("owner", "repo", 1)
assert "Could not retrieve comments" in res
assert "owner/repo" in res
def test_list_assigned_pull_requests_success() -> None:
mock_client = _create_mock_client()
repo: RepositoryModel = RepositoryModel(name="repo1", owner="owner1")
pr: PullRequestModel = PullRequestModel(number=1, title="Test PR")
mock_client.repos.list_all_user_repos.return_value = [repo]
mock_client.prs.list_assigned_pull_requests.return_value = [pr]
pr_tools: PRTools = PRTools(mock_client)
res: list[dict[str, Any]] = pr_tools.list_assigned_pull_requests()
assert len(res) == 1
assert res[0]["number"] == 1
mock_client.repos.list_all_user_repos.assert_called_once()
mock_client.prs.list_assigned_pull_requests.assert_called_once_with(
"owner1", "repo1"
)
def test_list_assigned_pull_requests_failure() -> None:
mock_client = _create_mock_client()
mock_client.repos.list_all_user_repos.side_effect = Exception("API Error")
pr_tools: PRTools = PRTools(mock_client)
res: list[dict[str, Any]] = pr_tools.list_assigned_pull_requests()
assert res == []
def test_list_pull_requests_success() -> None:
mock_client = _create_mock_client()
pr: PullRequestModel = PullRequestModel(number=1, title="Test PR")
mock_client.prs.list_repo_pull_requests.return_value = [pr]
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.list_pull_requests("owner", "repo")
assert "#1: Test PR" in res
def test_list_pull_requests_empty() -> None:
mock_client = _create_mock_client()
mock_client.prs.list_repo_pull_requests.return_value = []
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.list_pull_requests("owner", "repo")
assert res == "No open PRs in owner/repo."
def test_list_pull_requests_failure() -> None:
mock_client = _create_mock_client()
mock_client.prs.list_repo_pull_requests.side_effect = Exception("API Error")
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.list_pull_requests("owner", "repo")
assert "Could not list PRs" in res
assert "owner/repo" in res
def test_create_pull_request_success() -> None:
mock_client = _create_mock_client()
pr: PullRequestModel = PullRequestModel(number=2, title="Title")
mock_client.prs.create_pr_via_tea.return_value = pr
pr_tools: PRTools = PRTools(mock_client)
res: PullRequestModel = pr_tools.create_pull_request(
"owner", "repo", "head", "base", "Title", "Desc"
)
assert isinstance(res, PullRequestModel)
assert res.number == 2
mock_client.prs.create_pr_via_tea.assert_called_once_with(
"owner", "repo", "Title", "Desc", "head", "base"
)
def test_create_pull_request_failure() -> None:
mock_client = _create_mock_client()
mock_client.prs.create_pr_via_tea.side_effect = Exception("API Error")
pr_tools: PRTools = PRTools(mock_client)
try:
pr_tools.create_pull_request("owner", "repo", "head", "base", "Title")
assert False, "Expected Exception"
except Exception as e:
assert str(e) == "API Error"
def test_add_label_to_pr_success() -> None:
mock_client = _create_mock_client()
mock_client.prs.add_label_pr.return_value = {}
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.add_label_to_pr("owner", "repo", 1, "bug")
assert res == "Label 'bug' added to PR #1."
def test_add_label_to_pr_failure() -> None:
mock_client = _create_mock_client()
mock_client.prs.add_label_pr.side_effect = Exception("API Error")
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.add_label_to_pr("owner", "repo", 1, "bug")
assert "Could not add label" in res
assert "bug" in res
def test_get_pull_request_diff_success() -> None:
mock_client = _create_mock_client()
mock_client.prs.get_pull_request_diff.return_value = "diff content"
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.get_pull_request_diff("owner", "repo", 1)
assert res == "diff content"
def test_get_pull_request_diff_failure() -> None:
mock_client = _create_mock_client()
mock_client.prs.get_pull_request_diff.side_effect = Exception("API Error")
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.get_pull_request_diff("owner", "repo", 1)
assert "Could not retrieve diff" in res
assert "owner/repo" in res
def test_get_pull_request_patch_success() -> None:
mock_client = _create_mock_client()
mock_client.prs.get_pull_request_patch.return_value = "patch content"
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.get_pull_request_patch("owner", "repo", 1)
assert res == "patch content"
def test_get_pull_request_patch_failure() -> None:
mock_client = _create_mock_client()
mock_client.prs.get_pull_request_patch.side_effect = Exception("API Error")
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.get_pull_request_patch("owner", "repo", 1)
assert "Could not retrieve patch" in res
assert "owner/repo" in res
def test_approve_pull_request_success() -> None:
mock_client = _create_mock_client()
mock_client.prs.approve_pr.return_value = {}
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.approve_pull_request("owner", "repo", 1, "good")
assert res == "Approved PR #1."
def test_approve_pull_request_failure() -> None:
mock_client = _create_mock_client()
mock_client.prs.approve_pr.side_effect = Exception("API Error")
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.approve_pull_request("owner", "repo", 1, "good")
assert "Could not approve PR" in res
assert "owner/repo" in res
def test_request_changes_success() -> None:
mock_client = _create_mock_client()
mock_client.prs.request_changes_pr.return_value = {}
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.request_changes("owner", "repo", 1, "bad")
assert res == "Requested changes on PR #1."
def test_request_changes_failure() -> None:
mock_client = _create_mock_client()
mock_client.prs.request_changes_pr.side_effect = Exception("API Error")
pr_tools: PRTools = PRTools(mock_client)
res: str = pr_tools.request_changes("owner", "repo", 1, "bad")
assert "Could not request changes" in res
assert "owner/repo" in res