diff --git a/src/forge/orchestrator/worker.py b/src/forge/orchestrator/worker.py index 2a425870..6193f3b1 100644 --- a/src/forge/orchestrator/worker.py +++ b/src/forge/orchestrator/worker.py @@ -29,6 +29,13 @@ from forge.skills.utils import extract_project_key from forge.utils.redaction import redact_secrets from forge.workflow.nodes.error_handler import notify_error +from forge.workflow.pr_state import ( + activate_pull_request_for_event, + all_pull_requests_merged, + event_targets_pull_request, + mark_active_pull_request_merged, + save_active_pull_request, +) from forge.workflow.registry import create_default_router from forge.workflow.router import WorkflowRouter from forge.workflow.utils.automated_review_triage import ( @@ -439,6 +446,15 @@ async def _process_workflow(self, message: QueueMessage) -> None: # Run the workflow from the beginning result = await compiled_workflow.ainvoke(state, config=config) + # Nodes continue to use scalar PR fields as a compatibility view. + # Persist that view back into the selected per-repository record + # after every invocation so subsequent webhooks restore fresh CI, + # review, and merge state for the PR they target. + persisted_result = save_active_pull_request(result) + if persisted_result != result: + await compiled_workflow.aupdate_state(config, persisted_result) + result = persisted_result + final_node = result.get("current_node", "unknown") is_paused = result.get("is_paused", False) logger.info( @@ -480,6 +496,8 @@ async def _handle_resume_event( Updated state for workflow resumption. """ payload = message.payload + current_state = activate_pull_request_for_event(current_state, payload) + targets_implementation_pr = event_targets_pull_request(current_state, payload) changelog = payload.get("changelog", {}) comment = payload.get("comment", {}) @@ -499,6 +517,7 @@ async def _handle_resume_event( automated_review_revision_pending = None proposal_review_threads: list[dict[str, Any]] = [] proposal_review_decisions: list[dict[str, Any]] = [] + implementation_pr_approved = False current_node = current_state.get("current_node", "") @@ -556,12 +575,12 @@ async def _handle_resume_event( # GitHub fires check_suite webhooks for created/in_progress/completed — evaluating # on the earlier actions would see a partial set of check runs and could # prematurely declare success. Other event types (push, pull_request) always wake up. - if ( + event = message.event_type + is_check_event = "check_suite" in event or "check_run" in event + if message.source == EventSource.GITHUB and ( current_node in ("wait_for_ci_gate", "ci_evaluator") - and message.source == EventSource.GITHUB + or (targets_implementation_pr and is_check_event) ): - event = message.event_type - is_check_event = "check_suite" in event or "check_run" in event if is_check_event: suite_status = payload.get("check_suite", {}).get("status") or payload.get( "check_run", {} @@ -574,7 +593,11 @@ async def _handle_resume_event( else: is_ci_webhook = True logger.info(f"Detected GitHub CI webhook signal for {current_node}") - elif "issue_comment" not in event: + elif ( + "issue_comment" not in event + and "pull_request_review" not in event + and payload.get("pull_request", {}).get("merged") is not True + ): is_ci_webhook = True logger.info(f"Detected GitHub CI webhook signal for {current_node}") @@ -1320,7 +1343,7 @@ async def _handle_resume_event( if ( message.source == EventSource.GITHUB and "pull_request_review" in message.event_type - and current_node in _REVIEW_GATES + and (current_node in _REVIEW_GATES or targets_implementation_pr) and current_state.get("is_paused", True) ): review = payload.get("review", {}) @@ -1328,7 +1351,8 @@ async def _handle_resume_event( review_body = review.get("body", "") or "" if review_state == "approved": - # PR approved — advance to complete_tasks (via route_human_review default) + if targets_implementation_pr: + implementation_pr_approved = True is_approved = True logger.info(f"Detected PR review approval for {message.ticket_key}") elif review_state in ("changes_requested", "commented"): @@ -1383,7 +1407,7 @@ async def _handle_resume_event( message.source == EventSource.GITHUB and "pull_request" in message.event_type and payload.get("pull_request", {}).get("merged") is True - and current_node in _REVIEW_GATES + and (current_node in _REVIEW_GATES or targets_implementation_pr) ): is_approved = True pr_merged = True @@ -1401,6 +1425,12 @@ async def _handle_resume_event( "payload": payload, }, } + if targets_implementation_pr and is_ci_webhook: + updated_state["current_node"] = "ci_evaluator" + elif targets_implementation_pr and ( + "pull_request_review" in message.event_type or pr_merged + ): + updated_state["current_node"] = "human_review_gate" was_errored = _is_workflow_errored(current_state) @@ -1506,12 +1536,19 @@ async def _handle_resume_event( updated_state["feedback_comment"] = None updated_state["last_error"] = None elif is_approved: - updated_state["is_paused"] = False + updated_state["is_paused"] = implementation_pr_approved updated_state["revision_requested"] = False updated_state["feedback_comment"] = None updated_state["last_error"] = None + if implementation_pr_approved: + updated_state["human_review_status"] = "approved" if pr_merged: updated_state["pr_merged"] = True + if event_targets_pull_request(updated_state, payload): + updated_state = mark_active_pull_request_merged(updated_state) + updated_state["pr_merged"] = all_pull_requests_merged(updated_state) + if not updated_state["pr_merged"]: + updated_state["is_paused"] = True if is_prd_review: # Specification review is a separate artifact cycle and must # receive its own automated revision budget. @@ -1634,7 +1671,7 @@ async def _handle_resume_event( ) return current_state - return updated_state + return save_active_pull_request(updated_state) async def _post_resume_ack_comment( self, diff --git a/src/forge/workflow/base.py b/src/forge/workflow/base.py index e0b21a9c..fcc4776c 100644 --- a/src/forge/workflow/base.py +++ b/src/forge/workflow/base.py @@ -9,6 +9,7 @@ from langgraph.graph.message import add_messages from forge.models.workflow import TicketType +from forge.workflow.pr_state import PullRequestState class BaseState(TypedDict, total=False): @@ -47,6 +48,7 @@ class PRIntegrationState(TypedDict, total=False): workspace_path: str | None pr_urls: list[str] + pull_requests: dict[str, PullRequestState] current_pr_url: str | None current_pr_number: int | None current_repo: str | None diff --git a/src/forge/workflow/bug/state.py b/src/forge/workflow/bug/state.py index 2239d73f..91983f1b 100644 --- a/src/forge/workflow/bug/state.py +++ b/src/forge/workflow/bug/state.py @@ -81,6 +81,7 @@ def create_initial_bug_state(ticket_key: str, **kwargs: Any) -> BugState: "bug_fix_implemented": False, "workspace_path": None, "pr_urls": [], + "pull_requests": {}, "fork_owner": None, "fork_repo": None, "merge_conflicts": [], diff --git a/src/forge/workflow/feature/state.py b/src/forge/workflow/feature/state.py index f0fcf37a..f48c02aa 100644 --- a/src/forge/workflow/feature/state.py +++ b/src/forge/workflow/feature/state.py @@ -94,6 +94,7 @@ def create_initial_feature_state(ticket_key: str, **kwargs: Any) -> FeatureState "tasks_by_repo": {}, "workspace_path": None, "pr_urls": [], + "pull_requests": {}, "fork_owner": None, "fork_repo": None, "merge_conflicts": [], diff --git a/src/forge/workflow/nodes/ci_evaluator.py b/src/forge/workflow/nodes/ci_evaluator.py index 51bdac6f..30dc7bb4 100644 --- a/src/forge/workflow/nodes/ci_evaluator.py +++ b/src/forge/workflow/nodes/ci_evaluator.py @@ -45,7 +45,27 @@ async def evaluate_ci_status(state: WorkflowState) -> WorkflowState: Updated state with ci_status. """ ticket_key = state["ticket_key"] - pr_urls = state.get("pr_urls", []) + current_pr_url = state.get("current_pr_url") + pull_requests = state.get("pull_requests", {}) + active_pr = pull_requests.get(state.get("current_repo", "")) + if pull_requests: + if not ( + isinstance(active_pr, dict) + and active_pr.get("number") == state.get("current_pr_number") + and current_pr_url + ): + logger.error("Cannot evaluate CI: active PR state is inconsistent for %s", ticket_key) + return update_state_timestamp( + { + **state, + "ci_status": "failed", + "current_node": "ci_evaluator", + "last_error": "Active pull request state is inconsistent", + } + ) + pr_urls = [current_pr_url] + else: + pr_urls = state.get("pr_urls", []) ci_fix_attempt = state.get("ci_fix_attempt", 0) ci_fix_max = state.get("ci_fix_max_attempts", 5) settings = get_settings() diff --git a/src/forge/workflow/nodes/pr_creation.py b/src/forge/workflow/nodes/pr_creation.py index 769c3c81..f0f89c97 100644 --- a/src/forge/workflow/nodes/pr_creation.py +++ b/src/forge/workflow/nodes/pr_creation.py @@ -15,6 +15,7 @@ from forge.prompts import load_prompt from forge.workflow.nodes.code_review import sync_pr_description from forge.workflow.nodes.post_merge_summary import _extract_impact +from forge.workflow.pr_state import save_active_pull_request from forge.workflow.utils import update_state_timestamp from forge.workflow.utils.jira_status import post_status_comment from forge.workspace.git_ops import GitOperations @@ -329,16 +330,18 @@ async def create_pull_request(state: WorkflowState) -> WorkflowState: logger.warning(f"Failed to append review exhaustion section: {e}") return update_state_timestamp( - { - **state, - "pr_urls": pr_urls, - "current_pr_url": pr_url, - "current_pr_number": pr_number, - "fork_owner": pr_target.fork_owner, - "fork_repo": pr_target.fork_repo, - "current_node": "teardown_workspace", - "last_error": None, - } + save_active_pull_request( + { + **state, + "pr_urls": pr_urls, + "current_pr_url": pr_url, + "current_pr_number": pr_number, + "fork_owner": pr_target.fork_owner, + "fork_repo": pr_target.fork_repo, + "current_node": "teardown_workspace", + "last_error": None, + } + ) ) except Exception as e: diff --git a/src/forge/workflow/pr_state.py b/src/forge/workflow/pr_state.py new file mode 100644 index 00000000..8f205229 --- /dev/null +++ b/src/forge/workflow/pr_state.py @@ -0,0 +1,160 @@ +"""Per-pull-request lifecycle state for multi-repository workflows.""" + +from copy import deepcopy +from typing import Any, TypedDict + + +class PullRequestState(TypedDict, total=False): + url: str + number: int | None + repo: str + fork_owner: str | None + fork_repo: str | None + ci_status: str | None + ci_failed_checks: list[dict[str, Any]] + ci_skipped_checks: list[str] + ci_fix_attempt: int + human_review_status: str | None + review_comments: list[dict[str, Any]] + contested_comments: list[dict[str, Any]] + review_response_posted: bool + merged: bool + lifecycle_node: str + + +_ACTIVE_FIELDS = { + "current_pr_url": "url", + "current_pr_number": "number", + "fork_owner": "fork_owner", + "fork_repo": "fork_repo", + "ci_status": "ci_status", + "ci_failed_checks": "ci_failed_checks", + "ci_skipped_checks": "ci_skipped_checks", + "ci_fix_attempt": "ci_fix_attempt", + "human_review_status": "human_review_status", + "review_comments": "review_comments", + "contested_comments": "contested_comments", + "review_response_posted": "review_response_posted", +} + +_PR_LIFECYCLE_NODES = { + "wait_for_ci_gate", + "ci_evaluator", + "attempt_ci_fix", + "human_review_gate", + "implement_review", + "review_response_gate", + "rebase_pr", +} + + +def _event_pr_number(payload: dict[str, Any]) -> int | None: + number = payload.get("pull_request", {}).get("number") + if number is None: + number = payload.get("issue", {}).get("number") + if isinstance(number, int): + return number + for container in (payload.get("check_suite", {}), payload.get("check_run", {})): + pull_requests = container.get("pull_requests", []) + if pull_requests: + number = pull_requests[0].get("number") + return number if isinstance(number, int) else None + suite_pull_requests = container.get("check_suite", {}).get("pull_requests", []) + if suite_pull_requests: + number = suite_pull_requests[0].get("number") + return number if isinstance(number, int) else None + return None + + +def _event_pr_url(payload: dict[str, Any]) -> str | None: + url = payload.get("pull_request", {}).get("html_url") + return url if isinstance(url, str) and url else None + + +def _record_matches_event(record: dict[str, Any], payload: dict[str, Any]) -> bool: + number = _event_pr_number(payload) + if number is None: + return False + if record.get("number") == number: + return True + return record.get("number") is None and _event_pr_url(payload) == record.get("url") + + +def save_active_pull_request(state: dict[str, Any]) -> dict[str, Any]: + """Copy the scalar compatibility view into its per-repository PR record.""" + repo = state.get("current_repo") + if not repo or (state.get("current_pr_number") is None and not state.get("current_pr_url")): + return state + + existing_pull_requests = state.get("pull_requests", {}) + pull_requests = dict(existing_pull_requests) + record = deepcopy(existing_pull_requests.get(repo, {})) + for scalar, per_pr in _ACTIVE_FIELDS.items(): + if scalar in state: + record[per_pr] = state[scalar] + record["repo"] = repo + current_node = state.get("current_node") + if current_node in _PR_LIFECYCLE_NODES: + record["lifecycle_node"] = current_node + else: + record.setdefault("lifecycle_node", "wait_for_ci_gate") + pull_requests[repo] = record + return {**state, "pull_requests": pull_requests} + + +def activate_pull_request_for_event( + state: dict[str, Any], payload: dict[str, Any] +) -> dict[str, Any]: + """Select the PR targeted by a GitHub webhook as the scalar compatibility view.""" + repo = payload.get("repository", {}).get("full_name") + number = _event_pr_number(payload) + pull_requests = state.get("pull_requests", {}) + record = pull_requests.get(repo) if repo else None + if not isinstance(record, dict) or not _record_matches_event(record, payload): + return state + + activated = {**state, "current_repo": repo} + if record.get("number") is None: + updated_pull_requests = deepcopy(pull_requests) + updated_pull_requests[repo]["number"] = number + activated["pull_requests"] = updated_pull_requests + record = updated_pull_requests[repo] + if state.get("current_repo") != repo: + activated["workspace_path"] = None + for scalar, per_pr in _ACTIVE_FIELDS.items(): + if per_pr in record: + activated[scalar] = deepcopy(record[per_pr]) + # The webhook is authoritative for the number. Set it after restoring the + # compatibility fields so correctness does not depend on dict iteration order. + activated["current_pr_number"] = number + if record.get("lifecycle_node") in _PR_LIFECYCLE_NODES: + activated["current_node"] = record["lifecycle_node"] + return activated + + +def all_pull_requests_merged(state: dict[str, Any]) -> bool: + """Return true only when every created implementation PR has merged.""" + pull_requests = state.get("pull_requests", {}) + return bool(pull_requests) and all( + record.get("merged", False) for record in pull_requests.values() + ) + + +def mark_active_pull_request_merged(state: dict[str, Any]) -> dict[str, Any]: + """Mark the selected per-repository PR merged without changing other records.""" + repo = state.get("current_repo") + existing_pull_requests = state.get("pull_requests", {}) + pull_requests = dict(existing_pull_requests) + record = deepcopy(existing_pull_requests.get(repo)) if repo else None + if not isinstance(record, dict): + return state + record["merged"] = True + pull_requests[repo] = record + return {**state, "pull_requests": pull_requests} + + +def event_targets_pull_request(state: dict[str, Any], payload: dict[str, Any]) -> bool: + """Return whether a webhook identifies one of the implementation PR records.""" + repo = payload.get("repository", {}).get("full_name") + record = state.get("pull_requests", {}).get(repo) + return isinstance(record, dict) and _record_matches_event(record, payload) diff --git a/src/forge/workflow/task_takeover/state.py b/src/forge/workflow/task_takeover/state.py index eed2d289..e4405485 100644 --- a/src/forge/workflow/task_takeover/state.py +++ b/src/forge/workflow/task_takeover/state.py @@ -51,6 +51,7 @@ def create_initial_task_takeover_state(ticket_key: str, **kwargs: Any) -> TaskTa "updated_at": now, "workspace_path": None, "pr_urls": [], + "pull_requests": {}, "fork_owner": None, "fork_repo": None, "merge_conflicts": [], diff --git a/tests/unit/orchestrator/test_worker.py b/tests/unit/orchestrator/test_worker.py index ac68b0b3..defe7816 100644 --- a/tests/unit/orchestrator/test_worker.py +++ b/tests/unit/orchestrator/test_worker.py @@ -74,6 +74,153 @@ async def test_report_new_workflow_error_skips_non_reportable_errors( notify.assert_not_awaited() +def _multi_repo_pr_state() -> dict: + return { + "ticket_key": "TEST-123", + "current_node": "human_review_gate", + "is_paused": True, + "context": {}, + "current_repo": "acme/frontend", + "current_pr_number": 20, + "current_pr_url": "https://github.com/acme/frontend/pull/20", + "pr_merged": False, + "pull_requests": { + "acme/backend": { + "repo": "acme/backend", + "number": 10, + "url": "https://github.com/acme/backend/pull/10", + "merged": False, + "ci_status": "pending", + }, + "acme/frontend": { + "repo": "acme/frontend", + "number": 20, + "url": "https://github.com/acme/frontend/pull/20", + "merged": False, + "ci_status": "passed", + }, + }, + } + + +@pytest.mark.asyncio +async def test_multi_repo_merge_waits_for_every_pr() -> None: + worker = OrchestratorWorker(consumer_name="test-worker") + state = _multi_repo_pr_state() + + def merge_message(repo: str, number: int) -> QueueMessage: + return QueueMessage( + message_id=f"msg-{number}", + event_id=f"evt-{number}", + source=EventSource.GITHUB, + event_type="pull_request", + ticket_key="TEST-123", + payload={ + "action": "closed", + "pull_request": {"merged": True, "number": number}, + "repository": {"full_name": repo}, + }, + ) + + partial = await worker._handle_resume_event(merge_message("acme/backend", 10), state) + + assert partial["current_repo"] == "acme/backend" + assert partial["pull_requests"]["acme/backend"]["merged"] is True + assert partial["pull_requests"]["acme/frontend"]["merged"] is False + assert partial["pr_merged"] is False + assert partial["is_paused"] is True + + complete = await worker._handle_resume_event(merge_message("acme/frontend", 20), partial) + + assert complete["pr_merged"] is True + assert complete["is_paused"] is False + + +@pytest.mark.asyncio +async def test_multi_repo_ci_webhook_selects_earlier_pr_from_review_gate() -> None: + worker = OrchestratorWorker(consumer_name="test-worker") + message = QueueMessage( + message_id="msg-ci", + event_id="evt-ci", + source=EventSource.GITHUB, + event_type="check_suite", + ticket_key="TEST-123", + payload={ + "check_suite": {"status": "completed", "pull_requests": [{"number": 10}]}, + "repository": {"full_name": "acme/backend"}, + }, + ) + + result = await worker._handle_resume_event(message, _multi_repo_pr_state()) + + assert result["current_repo"] == "acme/backend" + assert result["current_pr_number"] == 10 + assert result["current_node"] == "ci_evaluator" + assert result["is_paused"] is False + + +@pytest.mark.asyncio +async def test_multi_repo_approval_uses_common_state_cleanup_path() -> None: + worker = OrchestratorWorker(consumer_name="test-worker") + state = _multi_repo_pr_state() + state["last_error"] = "stale review failure" + state["revision_requested"] = True + state["feedback_comment"] = "old feedback" + message = QueueMessage( + message_id="msg-approved", + event_id="evt-approved", + source=EventSource.GITHUB, + event_type="pull_request_review", + ticket_key="TEST-123", + payload={ + "review": {"state": "approved", "body": "Looks good"}, + "pull_request": {"number": 10}, + "repository": {"full_name": "acme/backend"}, + }, + ) + + result = await worker._handle_resume_event(message, state) + + assert result["current_repo"] == "acme/backend" + assert result["is_paused"] is True + assert result["last_error"] is None + assert result["revision_requested"] is False + assert result["feedback_comment"] is None + assert result["human_review_status"] == "approved" + assert result["pull_requests"]["acme/backend"]["human_review_status"] == "approved" + + +@pytest.mark.asyncio +@patch("forge.orchestrator.worker.GitHubClient") +async def test_multi_repo_review_selects_earlier_pr(mock_github_client: MagicMock) -> None: + github = AsyncMock() + github.get_review_comments.return_value = [] + mock_github_client.return_value = github + worker = OrchestratorWorker(consumer_name="test-worker") + state = _multi_repo_pr_state() + state["current_node"] = "wait_for_ci_gate" + message = QueueMessage( + message_id="msg-review", + event_id="evt-review", + source=EventSource.GITHUB, + event_type="pull_request_review", + ticket_key="TEST-123", + payload={ + "review": {"id": 5, "state": "changes_requested", "body": "Fix backend"}, + "pull_request": {"number": 10}, + "repository": {"full_name": "acme/backend"}, + }, + ) + + result = await worker._handle_resume_event(message, state) + + assert result["current_repo"] == "acme/backend" + assert result["current_pr_number"] == 10 + assert result["current_node"] == "human_review_gate" + assert result["revision_requested"] is True + assert result["feedback_comment"] == "Fix backend" + + class TestQuestionDetection: """Tests for Q&A mode question detection.""" diff --git a/tests/unit/workflow/nodes/test_ci_attempt_tracking.py b/tests/unit/workflow/nodes/test_ci_attempt_tracking.py index d8ce68eb..9cad5d96 100644 --- a/tests/unit/workflow/nodes/test_ci_attempt_tracking.py +++ b/tests/unit/workflow/nodes/test_ci_attempt_tracking.py @@ -36,6 +36,28 @@ def create_base_state(**kwargs) -> FeatureState: return FeatureState(**defaults) +@pytest.mark.asyncio +async def test_multi_pr_ci_rejects_inconsistent_active_view() -> None: + state = create_base_state( + current_repo="org/other", + current_pr_number=99, + current_pr_url="https://github.com/org/other/pull/99", + pull_requests={ + "org/repo": { + "repo": "org/repo", + "number": 42, + "url": "https://github.com/org/repo/pull/42", + } + }, + ) + + result = await evaluate_ci_status(state) + + assert result["ci_status"] == "failed" + assert result["current_node"] == "ci_evaluator" + assert result["last_error"] == "Active pull request state is inconsistent" + + # ── State Initialization Tests ──────────────────────────────────────────────── @@ -45,22 +67,26 @@ class TestCIAttemptTrackingStateFields: def test_current_attempt_in_ci_integration_state(self): """current_attempt must be a field in CIIntegrationState.""" from forge.workflow.base import CIIntegrationState + assert "ci_fix_attempt" in CIIntegrationState.__annotations__ def test_max_attempts_in_ci_integration_state(self): """max_attempts must be a field in CIIntegrationState.""" from forge.workflow.base import CIIntegrationState + assert "ci_fix_max_attempts" in CIIntegrationState.__annotations__ def test_feature_state_initializes_current_attempt_to_zero(self): """Feature state should initialize current_attempt to 0.""" from forge.workflow.feature.state import create_initial_feature_state + state = create_initial_feature_state(ticket_key="TEST-1") assert state.get("ci_fix_attempt") == 0 def test_feature_state_initializes_max_attempts_from_config(self): """Feature state should initialize max_attempts from config.""" from forge.workflow.feature.state import create_initial_feature_state + state = create_initial_feature_state(ticket_key="TEST-1") # Default config value is 5 assert state.get("ci_fix_max_attempts") is not None @@ -69,12 +95,14 @@ def test_feature_state_initializes_max_attempts_from_config(self): def test_bug_state_initializes_current_attempt_to_zero(self): """Bug state should initialize current_attempt to 0.""" from forge.workflow.bug.state import create_initial_bug_state + state = create_initial_bug_state(ticket_key="TEST-2") assert state.get("ci_fix_attempt") == 0 def test_bug_state_initializes_max_attempts_from_config(self): """Bug state should initialize max_attempts from config.""" from forge.workflow.bug.state import create_initial_bug_state + state = create_initial_bug_state(ticket_key="TEST-2") # Default config value is 5 assert state.get("ci_fix_max_attempts") is not None @@ -91,7 +119,7 @@ class TestCIAttemptIncrement: async def test_first_ci_failure_increments_attempt_to_one(self): """First CI failure should increment current_attempt from 0 to 1.""" state = create_base_state(ci_fix_attempt=0, ci_fix_max_attempts=3) - + github = create_mock_github_client() github.get_pull_request.return_value = {"head": {"sha": "abc123"}} github.get_check_runs.return_value = [ @@ -117,7 +145,7 @@ async def test_first_ci_failure_increments_attempt_to_one(self): async def test_second_ci_failure_increments_attempt_to_two(self): """Second CI failure should increment current_attempt from 1 to 2.""" state = create_base_state(ci_fix_attempt=1, ci_fix_max_attempts=3) - + github = create_mock_github_client() github.get_pull_request.return_value = {"head": {"sha": "abc123"}} github.get_check_runs.return_value = [ @@ -143,7 +171,7 @@ async def test_second_ci_failure_increments_attempt_to_two(self): async def test_third_ci_failure_increments_attempt_to_three(self): """Third CI failure should increment current_attempt from 2 to 3.""" state = create_base_state(ci_fix_attempt=2, ci_fix_max_attempts=3) - + github = create_mock_github_client() github.get_pull_request.return_value = {"head": {"sha": "abc123"}} github.get_check_runs.return_value = [ @@ -176,7 +204,7 @@ class TestCIAttemptLimitValidation: async def test_attempt_at_max_limit_blocks_further_attempts(self): """When current_attempt equals max_attempts, no more attempts should be made.""" state = create_base_state(ci_fix_attempt=3, ci_fix_max_attempts=3) - + github = create_mock_github_client() github.get_pull_request.return_value = {"head": {"sha": "abc123"}} github.get_check_runs.return_value = [ @@ -193,7 +221,9 @@ async def test_attempt_at_max_limit_blocks_further_attempts(self): with patch("forge.workflow.nodes.ci_evaluator.get_settings") as mock_settings: mock_settings.return_value.ci_fix_max_retries = 5 mock_settings.return_value.ignored_ci_checks = ["tide"] - with patch("forge.workflow.nodes.ci_evaluator.record_ci_fix_attempt") as mock_record: + with patch( + "forge.workflow.nodes.ci_evaluator.record_ci_fix_attempt" + ) as mock_record: result = await evaluate_ci_status(state) # Should not increment or route to attempt_ci_fix @@ -206,7 +236,7 @@ async def test_attempt_at_max_limit_blocks_further_attempts(self): async def test_attempt_exceeding_max_limit_blocks_further_attempts(self): """When current_attempt exceeds max_attempts, no more attempts should be made.""" state = create_base_state(ci_fix_attempt=4, ci_fix_max_attempts=3) - + github = create_mock_github_client() github.get_pull_request.return_value = {"head": {"sha": "abc123"}} github.get_check_runs.return_value = [ @@ -223,7 +253,9 @@ async def test_attempt_exceeding_max_limit_blocks_further_attempts(self): with patch("forge.workflow.nodes.ci_evaluator.get_settings") as mock_settings: mock_settings.return_value.ci_fix_max_retries = 5 mock_settings.return_value.ignored_ci_checks = ["tide"] - with patch("forge.workflow.nodes.ci_evaluator.record_ci_fix_attempt") as mock_record: + with patch( + "forge.workflow.nodes.ci_evaluator.record_ci_fix_attempt" + ) as mock_record: result = await evaluate_ci_status(state) # Should not increment or route to attempt_ci_fix @@ -236,7 +268,7 @@ async def test_attempt_exceeding_max_limit_blocks_further_attempts(self): async def test_attempt_one_below_max_allows_final_attempt(self): """When current_attempt is one below max, one more attempt should be allowed.""" state = create_base_state(ci_fix_attempt=2, ci_fix_max_attempts=3) - + github = create_mock_github_client() github.get_pull_request.return_value = {"head": {"sha": "abc123"}} github.get_check_runs.return_value = [ @@ -271,7 +303,7 @@ class TestCIAttemptReset: async def test_current_attempt_resets_on_ci_success(self): """When CI passes, current_attempt should reset to 0.""" state = create_base_state(ci_fix_attempt=2, ci_fix_max_attempts=3) - + github = create_mock_github_client() github.get_pull_request.return_value = {"head": {"sha": "abc123"}} github.get_check_runs.return_value = [ @@ -308,7 +340,7 @@ async def test_current_attempt_resets_on_ci_success(self): async def test_current_attempt_resets_on_workflow_completion(self): """When workflow completes (tasks complete), current_attempt should reset to 0.""" from forge.workflow.nodes.human_review import complete_tasks - + state = create_base_state( ci_fix_attempt=2, implemented_tasks=["TASK-1", "TASK-2"], @@ -339,7 +371,7 @@ async def test_missing_current_attempt_defaults_to_zero(self): state = create_base_state() # Remove current_attempt from state del state["ci_fix_attempt"] - + github = create_mock_github_client() github.get_pull_request.return_value = {"head": {"sha": "abc123"}} github.get_check_runs.return_value = [ @@ -367,7 +399,7 @@ async def test_missing_max_attempts_defaults_to_config_value(self): state = create_base_state(ci_fix_attempt=0) # Remove max_attempts from state del state["ci_fix_max_attempts"] - + github = create_mock_github_client() github.get_pull_request.return_value = {"head": {"sha": "abc123"}} github.get_check_runs.return_value = [ @@ -394,7 +426,7 @@ async def test_missing_max_attempts_defaults_to_config_value(self): async def test_max_attempts_one_allows_single_attempt(self): """When max_attempts is 1, only one attempt should be allowed.""" state = create_base_state(ci_fix_attempt=0, ci_fix_max_attempts=1) - + github = create_mock_github_client() github.get_pull_request.return_value = {"head": {"sha": "abc123"}} github.get_check_runs.return_value = [ diff --git a/tests/unit/workflow/nodes/test_pr_creation_pr_number.py b/tests/unit/workflow/nodes/test_pr_creation_pr_number.py index 40200b0a..2d691370 100644 --- a/tests/unit/workflow/nodes/test_pr_creation_pr_number.py +++ b/tests/unit/workflow/nodes/test_pr_creation_pr_number.py @@ -91,7 +91,9 @@ class TestPRNumberExtractionSuccess: @pytest.mark.asyncio async def test_pr_number_extracted_from_github_response(self): """Should extract PR number from GitHub API response and store in state.""" - mock_github = create_mock_github_client(pr_number=456, pr_url="https://github.com/owner/repo/pull/456") + mock_github = create_mock_github_client( + pr_number=456, pr_url="https://github.com/owner/repo/pull/456" + ) mock_jira = create_mock_jira_client() mock_git = create_mock_git_operations() @@ -107,8 +109,12 @@ async def test_pr_number_extracted_from_github_response(self): patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace()), - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, [])), + patch( + "forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace() + ), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, []) + ), patch("forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock), ): result = await create_pull_request(state) @@ -135,8 +141,12 @@ async def test_pr_number_used_in_jira_remote_link(self): patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace()), - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, [])), + patch( + "forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace() + ), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, []) + ), patch("forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock), ): await create_pull_request(state) @@ -165,8 +175,12 @@ async def test_pr_number_used_in_info_logging(self, caplog): patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace()), - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, [])), + patch( + "forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace() + ), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, []) + ), patch("forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock), ): await create_pull_request(state) @@ -202,14 +216,20 @@ async def test_pr_number_none_when_unavailable(self): patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace()), - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, [])), + patch( + "forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace() + ), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, []) + ), patch("forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock), ): result = await create_pull_request(state) # Verify PR number is None in state assert result["current_pr_number"] is None + assert result["pull_requests"]["owner/repo"]["number"] is None + assert result["pull_requests"]["owner/repo"]["url"] == result["current_pr_url"] @pytest.mark.asyncio async def test_workflow_continues_when_pr_number_unavailable(self): @@ -230,8 +250,12 @@ async def test_workflow_continues_when_pr_number_unavailable(self): patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace()), - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, [])), + patch( + "forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace() + ), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, []) + ), patch("forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock), ): result = await create_pull_request(state) @@ -264,8 +288,12 @@ async def test_warning_logged_when_pr_number_unavailable(self, caplog): patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace()), - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, [])), + patch( + "forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace() + ), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, []) + ), patch("forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock), ): await create_pull_request(state) @@ -298,8 +326,12 @@ async def test_generic_label_used_when_pr_number_unavailable(self): patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace()), - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, [])), + patch( + "forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace() + ), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, []) + ), patch("forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock), ): await create_pull_request(state) @@ -329,8 +361,12 @@ async def test_info_log_indicates_number_unavailable(self, caplog): patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace()), - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, [])), + patch( + "forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace() + ), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, []) + ), patch("forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock), ): await create_pull_request(state) @@ -338,8 +374,7 @@ async def test_info_log_indicates_number_unavailable(self, caplog): # Verify info log indicates number unavailable info_logs = [r for r in caplog.records if r.levelname == "INFO"] assert any( - "Created PR (number unavailable):" in record.message - and pr_url in record.message + "Created PR (number unavailable):" in record.message and pr_url in record.message for record in info_logs ) @@ -367,8 +402,12 @@ async def test_pr_number_zero_handled_correctly(self): patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace()), - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, [])), + patch( + "forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace() + ), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, []) + ), patch("forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock), ): result = await create_pull_request(state) @@ -401,8 +440,12 @@ async def test_pr_number_extracted_when_pr_url_missing(self): patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace()), - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, [])), + patch( + "forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace() + ), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, []) + ), patch("forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock), ): result = await create_pull_request(state) @@ -430,27 +473,41 @@ async def test_multiple_prs_each_have_own_pr_number(self): patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github_1), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace()), - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, [])), + patch( + "forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace() + ), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, []) + ), patch("forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock), ): result_1 = await create_pull_request(state) # Verify first PR has correct number assert result_1["current_pr_number"] == 100 + assert result_1["pull_requests"]["owner/repo"]["number"] == 100 # Simulate second PR creation with different number - mock_github_2 = create_mock_github_client(pr_number=200) + mock_github_2 = create_mock_github_client( + pr_number=200, pr_url="https://github.com/owner/other/pull/200" + ) + result_1["current_repo"] = "owner/other" with ( patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github_2), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace()), - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, [])), + patch( + "forge.workflow.nodes.pr_creation.Workspace", return_value=create_mock_workspace() + ), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", return_value=(False, []) + ), patch("forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock), ): result_2 = await create_pull_request(result_1) # Verify second PR has correct number assert result_2["current_pr_number"] == 200 + assert result_2["pull_requests"]["owner/repo"]["number"] == 100 + assert result_2["pull_requests"]["owner/other"]["number"] == 200 diff --git a/tests/unit/workflow/test_pr_state.py b/tests/unit/workflow/test_pr_state.py new file mode 100644 index 00000000..e3cf8892 --- /dev/null +++ b/tests/unit/workflow/test_pr_state.py @@ -0,0 +1,198 @@ +from forge.workflow.pr_state import ( + activate_pull_request_for_event, + all_pull_requests_merged, + event_targets_pull_request, + save_active_pull_request, +) + + +def _multi_repo_state() -> dict: + return { + "current_repo": "acme/frontend", + "current_pr_number": 20, + "current_pr_url": "https://github.com/acme/frontend/pull/20", + "fork_owner": "forge-bot", + "fork_repo": "frontend", + "ci_status": "passed", + "pr_merged": False, + "pull_requests": { + "acme/backend": { + "repo": "acme/backend", + "number": 10, + "url": "https://github.com/acme/backend/pull/10", + "fork_owner": "forge-bot", + "fork_repo": "backend", + "ci_status": "fixing", + "ci_failed_checks": [{"name": "unit"}], + "ci_fix_attempt": 2, + "merged": False, + "lifecycle_node": "review_response_gate", + }, + "acme/frontend": { + "repo": "acme/frontend", + "number": 20, + "url": "https://github.com/acme/frontend/pull/20", + "fork_owner": "forge-bot", + "fork_repo": "frontend", + "ci_status": "passed", + "merged": False, + "lifecycle_node": "human_review_gate", + }, + }, + } + + +def test_github_event_activates_matching_repo_pr() -> None: + state = _multi_repo_state() + payload = { + "repository": {"full_name": "acme/backend"}, + "pull_request": {"number": 10}, + } + + activated = activate_pull_request_for_event(state, payload) + + assert activated["current_repo"] == "acme/backend" + assert activated["current_pr_number"] == 10 + assert activated["fork_repo"] == "backend" + assert activated["ci_status"] == "fixing" + assert activated["ci_failed_checks"] == [{"name": "unit"}] + assert activated["current_node"] == "review_response_gate" + assert activated["workspace_path"] is None + + +def test_event_for_unknown_pr_does_not_change_active_repo() -> None: + state = _multi_repo_state() + payload = { + "repository": {"full_name": "acme/other"}, + "pull_request": {"number": 99}, + } + + assert activate_pull_request_for_event(state, payload) == state + + +def test_event_without_pr_number_does_not_target_record_without_number() -> None: + state = _multi_repo_state() + state["pull_requests"]["acme/backend"].pop("number") + payload = {"repository": {"full_name": "acme/backend"}} + + assert not event_targets_pull_request(state, payload) + assert activate_pull_request_for_event(state, payload) == state + + +def test_check_run_event_activates_matching_repo_pr() -> None: + state = _multi_repo_state() + payload = { + "repository": {"full_name": "acme/backend"}, + "check_run": {"pull_requests": [{"number": 10}]}, + } + + activated = activate_pull_request_for_event(state, payload) + + assert activated["current_repo"] == "acme/backend" + assert activated["current_pr_number"] == 10 + + +def test_same_repo_event_preserves_existing_workspace() -> None: + state = _multi_repo_state() + state["workspace_path"] = "/tmp/forge-AISOS-1-active" + payload = { + "repository": {"full_name": "acme/frontend"}, + "pull_request": {"number": 20}, + } + + activated = activate_pull_request_for_event(state, payload) + + assert activated["workspace_path"] == "/tmp/forge-AISOS-1-active" + + +def test_issue_comment_event_activates_matching_repo_pr() -> None: + state = _multi_repo_state() + payload = { + "repository": {"full_name": "acme/backend"}, + "issue": {"number": 10}, + } + + activated = activate_pull_request_for_event(state, payload) + + assert activated["current_repo"] == "acme/backend" + assert activated["current_pr_number"] == 10 + + +def test_active_changes_are_saved_only_to_matching_repo() -> None: + state = activate_pull_request_for_event( + _multi_repo_state(), + {"repository": {"full_name": "acme/backend"}, "pull_request": {"number": 10}}, + ) + state["ci_status"] = "passed" + state["ci_fix_attempt"] = 0 + + saved = save_active_pull_request(state) + + assert saved["pull_requests"]["acme/backend"]["ci_status"] == "passed" + assert saved["pull_requests"]["acme/backend"]["ci_fix_attempt"] == 0 + assert saved["pull_requests"]["acme/frontend"]["ci_status"] == "passed" + + +def test_merge_completion_requires_every_pr() -> None: + state = _multi_repo_state() + state["pull_requests"]["acme/backend"]["merged"] = True + assert not all_pull_requests_merged(state) + + state["pull_requests"]["acme/frontend"]["merged"] = True + assert all_pull_requests_merged(state) + + +def test_url_only_pr_blocks_aggregate_merge_completion() -> None: + state = _multi_repo_state() + state["pull_requests"]["acme/backend"]["merged"] = True + state["pull_requests"]["acme/frontend"]["merged"] = True + state["pull_requests"]["acme/docs"] = { + "repo": "acme/docs", + "url": "https://github.com/acme/docs/pull/30", + "number": None, + "merged": False, + } + + assert not all_pull_requests_merged(state) + + +def test_later_webhook_hydrates_url_only_pr_number() -> None: + state = _multi_repo_state() + state["pull_requests"]["acme/docs"] = { + "repo": "acme/docs", + "url": "https://github.com/acme/docs/pull/30", + "number": None, + "merged": False, + } + payload = { + "repository": {"full_name": "acme/docs"}, + "pull_request": { + "number": 30, + "html_url": "https://github.com/acme/docs/pull/30", + }, + } + + activated = activate_pull_request_for_event(state, payload) + + assert activated["current_repo"] == "acme/docs" + assert activated["current_pr_number"] == 30 + assert activated["pull_requests"]["acme/docs"]["number"] == 30 + + +def test_url_only_pr_rejects_different_pr_in_same_repo() -> None: + state = _multi_repo_state() + state["pull_requests"]["acme/docs"] = { + "repo": "acme/docs", + "url": "https://github.com/acme/docs/pull/30", + "number": None, + "merged": False, + } + payload = { + "repository": {"full_name": "acme/docs"}, + "pull_request": { + "number": 31, + "html_url": "https://github.com/acme/docs/pull/31", + }, + } + + assert not event_targets_pull_request(state, payload)