diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 316ee7f87..72be2c989 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -4,6 +4,8 @@ on: pull_request: branches: - main + - integration/dev-to-main + - prototype/declarative-workflows jobs: helm-lint: diff --git a/CLAUDE.md b/CLAUDE.md index 1eb1e8936..257e6852c 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -138,7 +138,12 @@ podman rm $(podman ps -a --filter name=forge- -q) | `/forge unskip-gate ` | PR comment | Remove a previously set skip | | `/forge rebase` | PR comment | Merge main into PR branch, resolving conflicts with AI | -Skip-gate commands are only active at CI stages (`wait_for_ci_gate`, `ci_evaluator`, `attempt_ci_fix`). Rebase works from any workflow stage. +Skip-gate commands are only active at CI stages (`ci_evaluator`, `attempt_ci_fix`, `human_review_gate`). Rebase works from any workflow stage. + +## Source Control Registry (repos.yaml) + +`load_registry()` parses the optional `repos.yaml`-shaped config referenced by `settings.forge_repos_config_path`, resolving repository identifiers to connections and adapters. `get_registry()` caches the result for the life of the process (`@lru_cache`) — editing `repos.yaml` requires a **process restart** to take effect; there is no runtime reload. A misconfigured `repos.yaml` (unknown provider, missing `credential_env`, etc.) fails app startup rather than the first inbound webhook. + ## PRD & Spec Approval via GitHub PR diff --git a/docs/roadmap.md b/docs/roadmap.md index 2ff7e9a41..5feda7e9c 100644 --- a/docs/roadmap.md +++ b/docs/roadmap.md @@ -1,7 +1,9 @@ # Forge Product Roadmap -**Status:** Draft for discussion -**Planning horizon:** Outcome-based; dates and release assignments are intentionally TBD +**Status:** Draft for discussion + +**Planning horizon:** Outcome-based; dates and release assignments are intentionally TBD + **Last reviewed:** 2026-08-10 ## Product direction diff --git a/src/forge/api/routes/github.py b/src/forge/api/routes/github.py index 44868faa4..fbb239e25 100644 --- a/src/forge/api/routes/github.py +++ b/src/forge/api/routes/github.py @@ -1,8 +1,9 @@ """GitHub webhook endpoint for receiving repository events.""" import hashlib -import hmac +import json import logging +import re from typing import Any from fastapi import APIRouter, Header, HTTPException, Request, status @@ -13,20 +14,56 @@ record_webhook_received, ) from forge.config import get_settings -from forge.integrations.github.webhooks import ( - create_github_webhook_event, - parse_github_webhook, -) -from forge.models.events import EventSource +from forge.integrations.source_control.contracts import Connection, NormalizedEvent, Provider +from forge.integrations.source_control.errors import NotFoundError, ProviderConfigError +from forge.integrations.source_control.github.adapter import GitHubAdapter +from forge.integrations.source_control.registry import get_registry, resolve_env_value from forge.observability.config import get_tracer from forge.observability.context import get_correlation_id from forge.queue.producer import QueueProducer +_DEFAULT_GITHUB_CONNECTION = Connection( + name="default-github", + provider=Provider.GITHUB, + base_url="https://api.github.com", + credential_env="GITHUB_TOKEN", + webhook_secret_env="GITHUB_WEBHOOK_SECRET", +) + logger = logging.getLogger(__name__) tracer = get_tracer("forge.api.github") router = APIRouter(prefix="/api/v1/webhooks", tags=["github"]) +TICKET_PATTERN = re.compile(r"([A-Z][A-Z0-9]+-\d+)", re.IGNORECASE) + + +def _extract_ticket_key(event: NormalizedEvent) -> str: + """Extract a Jira ticket key from a NormalizedEvent. + + Prefers the change request's title/branch when one is present (PR, review, + and check events carrying a pull_requests stub). Falls back to whatever + branch name the raw payload carries when there is no change request: + a push event's ref, or a check_suite/check_run event that fired before + GitHub attached a pull_requests stub to it (both still carry head_branch). + """ + if event.change_request is not None: + for text in (event.change_request.title, event.change_request.source_branch): + match = TICKET_PATTERN.search(text or "") + if match: + return match.group(1).upper() + raw = event.raw + branch_sources = ( + raw.get("ref", ""), + raw.get("check_suite", {}).get("head_branch", ""), + raw.get("check_run", {}).get("check_suite", {}).get("head_branch", ""), + ) + for text in branch_sources: + match = TICKET_PATTERN.search(str(text)) + if match: + return match.group(1).upper() + return "" + @router.post( "/github", @@ -47,7 +84,7 @@ async def receive_github_webhook( This endpoint: 1. Validates webhook signature - 2. Parses the webhook payload + 2. Parses the webhook payload into a NormalizedEvent 3. Queues the event for async processing 4. Returns immediately (<500ms target) @@ -80,9 +117,35 @@ async def receive_github_webhook( # Read raw body for signature verification body = await request.body() - # Validate signature - secret = settings.github_webhook_secret.get_secret_value() - if secret and not _verify_github_signature(body, x_hub_signature_256, secret): + try: + sniff_payload = json.loads(body) if body else {} + except json.JSONDecodeError: + sniff_payload = {} + repo_namespace = sniff_payload.get("repository", {}).get("full_name", "") + + registry = get_registry() + connection = _DEFAULT_GITHUB_CONNECTION + if repo_namespace: + try: + connection = registry.resolve( + repo_namespace, provider_hint=Provider.GITHUB + ).connection + except (NotFoundError, ProviderConfigError): + connection = _DEFAULT_GITHUB_CONNECTION + + webhook_secret = ( + resolve_env_value(connection.webhook_secret_env, settings) + if connection.webhook_secret_env + else None + ) + adapter = GitHubAdapter(connection=connection, webhook_secret=webhook_secret) + verify_headers = { + "X-GitHub-Event": x_github_event, + "X-GitHub-Delivery": x_github_delivery, + "X-Hub-Signature-256": x_hub_signature_256, + } + + if not await adapter.verify_webhook(verify_headers, body): span.set_attribute("error", True) span.set_attribute("error.type", "auth_failure") logger.warning("Invalid GitHub webhook signature") @@ -94,7 +157,7 @@ async def receive_github_webhook( # Parse JSON payload try: payload: dict[str, Any] = await request.json() - except Exception as e: + except (json.JSONDecodeError, UnicodeDecodeError) as e: span.set_attribute("error", True) span.set_attribute("error.type", "parse_error") logger.error(f"Failed to parse GitHub webhook payload: {e}") @@ -105,46 +168,43 @@ async def receive_github_webhook( event_id = x_github_delivery or _generate_event_id(payload) span.set_attribute("forge.event_id", event_id) + headers = {**verify_headers, "X-GitHub-Delivery": event_id} # Parse webhook data - webhook_data = parse_github_webhook(payload, x_github_event, event_id) + try: + event = await adapter.parse_webhook(headers, body, registry) + except (NotFoundError, ProviderConfigError): + span.set_attribute("forge.skipped", True) + span.set_attribute("forge.skip_reason", "unmanaged_repository") + logger.info("Skipping webhook for unmanaged repository") + record_webhook_received(source="github", event_type=x_github_event) + return {"status": "ignored", "event_id": event_id} + + ticket_key = _extract_ticket_key(event) + span.set_attribute("forge.ticket_key", ticket_key) # Record webhook received metric record_webhook_received(source="github", event_type=x_github_event) - span.set_attribute("forge.ticket_key", webhook_data.ticket_key or "") - webhook_event = create_github_webhook_event(webhook_data) - # Queue for async processing producer = QueueProducer() - message_id = await producer.publish_once( - event_id=webhook_event.event_id, - source=EventSource.GITHUB, - event_type=webhook_event.event_type, - ticket_key=webhook_event.ticket_key, - payload=webhook_event.payload, - ) + message_id = await producer.publish_event(event, ticket_key) if message_id is None: span.set_attribute("forge.skipped", True) span.set_attribute("forge.skip_reason", "duplicate event") - return { - "status": "duplicate", - "event_id": event_id, - "ticket_key": webhook_data.ticket_key or "", - } + return {"status": "duplicate", "event_id": event_id, "ticket_key": ticket_key} span.set_attribute("forge.queued", True) - logger.info(f"Queued GitHub event {event_id} for {webhook_data.ticket_key}") + logger.info( + f"GitHub webhook queued: event_id={event_id}, " + f"kind={event.kind}, repo={event.repo_ref.namespace}" + ) # Record webhook processed metric record_webhook_processed(source="github", event_type=x_github_event) - return { - "status": "accepted", - "event_id": event_id, - "ticket_key": webhook_data.ticket_key or "", - } + return {"status": "queued", "event_id": event_id, "ticket_key": ticket_key} except HTTPException: raise @@ -162,54 +222,28 @@ async def receive_github_webhook( except Exception as e: span.set_attribute("error", True) span.set_attribute("error.type", "internal_error") - logger.error(f"Failed to queue GitHub event: {e}") + logger.error(f"Failed to process GitHub webhook: {e}") record_webhook_failed( source="github", event_type=x_github_event, error_type="internal_error" ) raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail="Failed to queue event", + detail="Failed to process webhook", ) finally: span.end() -def _verify_github_signature(payload: bytes, signature: str, secret: str) -> bool: - """Verify GitHub webhook signature. - - Args: - payload: Raw request body. - signature: X-Hub-Signature-256 header value. - secret: Webhook secret. - - Returns: - True if signature is valid. - """ - if not signature: - return False - - expected = ( - "sha256=" - + hmac.new( - secret.encode("utf-8"), - payload, - hashlib.sha256, - ).hexdigest() - ) - - return hmac.compare_digest(signature, expected) - - def _generate_event_id(payload: dict[str, Any]) -> str: """Generate a deterministic event ID from payload. + Used as a fallback when X-GitHub-Delivery is empty. + Args: payload: Webhook payload. Returns: SHA256-based event ID. """ - import json - content = json.dumps(payload, sort_keys=True) return hashlib.sha256(content.encode()).hexdigest()[:16] diff --git a/src/forge/api/routes/health.py b/src/forge/api/routes/health.py index 14d7818e6..c92d402e2 100644 --- a/src/forge/api/routes/health.py +++ b/src/forge/api/routes/health.py @@ -8,6 +8,7 @@ from forge import __version__ from forge.orchestrator.checkpointer import get_redis_client +from forge.queue.producer import JIRA_STREAM, LEGACY_SOURCE_CONTROL_STREAM, SOURCE_CONTROL_STREAM logger = logging.getLogger(__name__) @@ -47,9 +48,10 @@ async def health_check() -> Any: # Get approximate queue depth try: - jira_len = await redis_client.xlen("forge:events:jira") - github_len = await redis_client.xlen("forge:events:github") - queue_depth = jira_len + github_len + jira_len = await redis_client.xlen(JIRA_STREAM) + source_control_len = await redis_client.xlen(SOURCE_CONTROL_STREAM) + legacy_len = await redis_client.xlen(LEGACY_SOURCE_CONTROL_STREAM) + queue_depth = jira_len + source_control_len + legacy_len except Exception: pass # Streams may not exist yet diff --git a/src/forge/cli.py b/src/forge/cli.py index 78cfb4d53..bad0eb5df 100644 --- a/src/forge/cli.py +++ b/src/forge/cli.py @@ -6,6 +6,7 @@ import sys from typing import Any +import forge.integrations.source_control.github # noqa: F401 (registers GitHub adapter factory) from forge.config import get_settings diff --git a/src/forge/config.py b/src/forge/config.py index 5d5572b26..9c8d2895d 100644 --- a/src/forge/config.py +++ b/src/forge/config.py @@ -81,7 +81,10 @@ def atlassian_auth_base64(self) -> str: # GitHub Configuration github_token: SecretStr = Field(description="GitHub personal access token") github_webhook_secret: SecretStr = Field( - default=SecretStr(""), description="Shared secret for GitHub webhook validation" + default=SecretStr(""), + description=( + "Shared secret for GitHub webhook validation; leave empty to accept unsigned webhooks" + ), ) github_default_repo: str = Field( default="", @@ -104,6 +107,14 @@ def atlassian_auth_base64(self) -> str: default="", description="GitHub account/org where forks are created (defaults to authenticated user if empty)", ) + forge_repos_config_path: str = Field( + default="config/repos.yaml", + description=( + "Path to the source control provider/connection registry config file. " + "Optional — repositories not listed here resolve through the implicit " + "default GitHub connection built from GITHUB_TOKEN/GITHUB_WEBHOOK_SECRET." + ), + ) forge_bot_comment_prefix: str = Field( default="", description="Prefix to use for all comments made by the Forge bot", diff --git a/src/forge/integrations/github/__init__.py b/src/forge/integrations/github/__init__.py index 7c6180bf5..2162e6ef7 100644 --- a/src/forge/integrations/github/__init__.py +++ b/src/forge/integrations/github/__init__.py @@ -1,25 +1,7 @@ """GitHub integration for PR management and webhook handling.""" from forge.integrations.github.client import GitHubClient -from forge.integrations.github.webhooks import ( - GitHubWebhookData, - create_github_webhook_event, - is_ci_failure, - is_ci_success, - is_pr_merged, - is_pr_review_approved, - is_pr_review_changes_requested, - parse_github_webhook, -) __all__ = [ "GitHubClient", - "GitHubWebhookData", - "create_github_webhook_event", - "is_ci_failure", - "is_ci_success", - "is_pr_merged", - "is_pr_review_approved", - "is_pr_review_changes_requested", - "parse_github_webhook", ] diff --git a/src/forge/integrations/github/client.py b/src/forge/integrations/github/client.py index 37b1a83e4..0e66ccbd7 100644 --- a/src/forge/integrations/github/client.py +++ b/src/forge/integrations/github/client.py @@ -53,14 +53,27 @@ class GitHubClient: pull requests, and code reviews. """ - def __init__(self, settings: Settings | None = None): + def __init__( + self, + settings: Settings | None = None, + *, + base_url: str | None = None, + ca_path: str | None = None, + ): """Initialize the GitHub client. Args: settings: Application settings. Uses default if not provided. + base_url: API base URL for the target connection (e.g. a GitHub + Enterprise Server's ``https://ghe.example.com/api/v3``). Falls + back to the public GitHub API when not provided. + ca_path: Path to a CA bundle to verify the connection's TLS + certificate against, for self-signed Enterprise Server + deployments. Falls back to the default trust store. """ self.settings = settings or get_settings() - self.base_url = "https://api.github.com" + self.base_url = base_url or "https://api.github.com" + self._ca_path = ca_path self._client: httpx.AsyncClient | None = None async def _get_client(self) -> httpx.AsyncClient: @@ -74,6 +87,7 @@ async def _get_client(self) -> httpx.AsyncClient: "X-GitHub-Api-Version": "2022-11-28", }, timeout=30.0, + verify=self._ca_path or True, ) return self._client @@ -171,6 +185,37 @@ async def get_pull_request(self, owner: str, repo: str, pr_number: int) -> dict[ response.raise_for_status() return response.json() + async def get_reviews(self, owner: str, repo: str, pr_number: int) -> list[dict[str, Any]]: + """Get the formal review submissions on a pull request. + + Each item carries a submission-level ``state`` (APPROVED, + CHANGES_REQUESTED, COMMENTED, DISMISSED, or PENDING), which is distinct + from the inline review comment threads returned by + get_pull_request_review_threads. + + Args: + owner: Repository owner. + repo: Repository name. + pr_number: Pull request number. + + Returns: + List of review submission dicts. + """ + client = await self._get_client() + reviews: list[dict[str, Any]] = [] + page = 1 + while True: + response = await client.get( + f"/repos/{owner}/{repo}/pulls/{pr_number}/reviews", + params={"per_page": 100, "page": page}, + ) + response.raise_for_status() + batch: list[dict[str, Any]] = response.json() + reviews.extend(batch) + if len(batch) < 100: + return reviews + page += 1 + async def create_review_comment( self, owner: str, @@ -513,6 +558,39 @@ async def get_workflow_run_logs(self, owner: str, repo: str, run_id: int) -> str response.raise_for_status() return response.text + async def list_workflow_run_jobs( + self, owner: str, repo: str, run_id: int | str + ) -> list[dict[str, Any]]: + """List the jobs belonging to a workflow run. + + This is how a check run is resolved to its Actions job: a check-run id is + a Checks API resource id, not an Actions job id, so callers match a job + by name within the run to recover the job id that /actions/jobs/{id}/logs + expects. + + Args: + owner: Repository owner. + repo: Repository name. + run_id: Workflow run ID. + + Returns: + List of job dicts, each carrying its Actions ``id`` and ``name``. The + default ``latest`` filter returns only the most recent attempt of + each job, so re-runs don't surface as duplicate same-named entries. + + Capped at the first 100 jobs (one page); Link-header pagination is + not followed, consistent with the other list endpoints on this + client. A workflow run with more than 100 jobs would truncate here, + in which case a job past the cap can't be resolved for log fetching. + """ + client = await self._get_client() + response = await client.get( + f"/repos/{owner}/{repo}/actions/runs/{run_id}/jobs", + params={"per_page": 100, "filter": "latest"}, + ) + response.raise_for_status() + return response.json().get("jobs", []) + async def get_job_logs(self, owner: str, repo: str, job_id: int | str) -> str: """Download logs for a specific Actions job as plain text. diff --git a/src/forge/integrations/github/webhooks.py b/src/forge/integrations/github/webhooks.py deleted file mode 100644 index eea1026e9..000000000 --- a/src/forge/integrations/github/webhooks.py +++ /dev/null @@ -1,265 +0,0 @@ -"""GitHub webhook payload parsing and validation.""" - -import logging -import re -from dataclasses import dataclass -from typing import Any - -from forge.models.events import EventSource, WebhookEvent - -logger = logging.getLogger(__name__) - -# Pattern to extract Jira ticket keys from PR titles/branches -TICKET_PATTERN = re.compile(r"([A-Z][A-Z0-9]+-\d+)", re.IGNORECASE) - - -@dataclass -class GitHubWebhookData: - """Parsed data from a GitHub webhook payload.""" - - event_id: str - event_type: str - action: str - repo_full_name: str - ticket_key: str | None - pr_number: int | None - pr_url: str | None - pr_state: str | None - branch_name: str | None - commit_sha: str | None - check_status: str | None - check_conclusion: str | None - sender_login: str - raw_payload: dict[str, Any] - - -def parse_github_webhook( - payload: dict[str, Any], - event_type: str, - event_id: str, -) -> GitHubWebhookData: - """Parse a GitHub webhook payload into structured data. - - Args: - payload: Raw webhook payload from GitHub. - event_type: Event type from X-GitHub-Event header. - event_id: Unique event identifier from X-GitHub-Delivery header. - - Returns: - Parsed GitHubWebhookData. - """ - action = payload.get("action", "") - repo = payload.get("repository", {}) - sender = payload.get("sender", {}) - - # Common fields - repo_full_name = repo.get("full_name", "") - sender_login = sender.get("login", "") - - # Extract ticket key from various sources - ticket_key = None - pr_number = None - pr_url = None - pr_state = None - branch_name = None - commit_sha = None - check_status = None - check_conclusion = None - - # Handle pull_request events - if event_type == "pull_request": - pr = payload.get("pull_request", {}) - pr_number = pr.get("number") - pr_url = pr.get("html_url") - pr_state = pr.get("state") - branch_name = pr.get("head", {}).get("ref", "") - - # Try to extract ticket from PR title first, then branch name - pr_title = pr.get("title", "") - ticket_key = _extract_ticket_key(pr_title) or _extract_ticket_key(branch_name) - - # Handle check_run events (CI status) - elif event_type == "check_run": - check_run = payload.get("check_run", {}) - check_status = check_run.get("status") - check_conclusion = check_run.get("conclusion") - commit_sha = check_run.get("head_sha") - - # Get ticket from associated PRs - pull_requests = check_run.get("pull_requests", []) - if pull_requests: - pr = pull_requests[0] - pr_number = pr.get("number") - pr_url = pr.get("url") - branch_name = pr.get("head", {}).get("ref", "") - ticket_key = _extract_ticket_key(branch_name) - - # Handle check_suite events - elif event_type == "check_suite": - check_suite = payload.get("check_suite", {}) - check_status = check_suite.get("status") - check_conclusion = check_suite.get("conclusion") - commit_sha = check_suite.get("head_sha") - branch_name = check_suite.get("head_branch", "") - ticket_key = _extract_ticket_key(branch_name) - - # Get ticket from associated PRs - pull_requests = check_suite.get("pull_requests", []) - if pull_requests: - pr = pull_requests[0] - pr_number = pr.get("number") - - # Handle pull request review submissions and inline thread replies - elif event_type in ("pull_request_review", "pull_request_review_comment"): - pr = payload.get("pull_request", {}) - pr_number = pr.get("number") - pr_url = pr.get("html_url") - pr_state = pr.get("state") - branch_name = pr.get("head", {}).get("ref", "") - pr_title = pr.get("title", "") - ticket_key = _extract_ticket_key(pr_title) or _extract_ticket_key(branch_name) - - # Handle push events - elif event_type == "push": - branch_name = payload.get("ref", "").replace("refs/heads/", "") - commit_sha = payload.get("after") - ticket_key = _extract_ticket_key(branch_name) - - # Handle issue_comment events (PR comments) - elif event_type == "issue_comment": - issue = payload.get("issue", {}) - # Check if this is a PR (has pull_request field) - if issue.get("pull_request"): - pr_number = issue.get("number") - pr_url = issue.get("html_url") - pr_title = issue.get("title", "") - ticket_key = _extract_ticket_key(pr_title) - - return GitHubWebhookData( - event_id=event_id, - event_type=event_type, - action=action, - repo_full_name=repo_full_name, - ticket_key=ticket_key, - pr_number=pr_number, - pr_url=pr_url, - pr_state=pr_state, - branch_name=branch_name, - commit_sha=commit_sha, - check_status=check_status, - check_conclusion=check_conclusion, - sender_login=sender_login, - raw_payload=payload, - ) - - -def create_github_webhook_event(data: GitHubWebhookData) -> WebhookEvent: - """Create a WebhookEvent from parsed GitHub webhook data. - - Args: - data: Parsed GitHub webhook data. - - Returns: - WebhookEvent ready for queue publishing. - """ - return WebhookEvent( - event_id=data.event_id, - source=EventSource.GITHUB, - event_type=f"{data.event_type}:{data.action}" if data.action else data.event_type, - ticket_key=data.ticket_key or "", - payload=data.raw_payload, - ) - - -def _extract_ticket_key(text: str) -> str | None: - """Extract Jira ticket key from text. - - Args: - text: Text that may contain a ticket key. - - Returns: - First ticket key found, or None. - """ - if not text: - return None - - match = TICKET_PATTERN.search(text) - if match: - return match.group(1).upper() - return None - - -def is_ci_success(data: GitHubWebhookData) -> bool: - """Check if the webhook indicates CI success. - - Args: - data: Parsed webhook data. - - Returns: - True if CI has passed. - """ - return data.check_status == "completed" and data.check_conclusion == "success" - - -def is_ci_failure(data: GitHubWebhookData) -> bool: - """Check if the webhook indicates CI failure. - - Args: - data: Parsed webhook data. - - Returns: - True if CI has failed. - """ - return data.check_status == "completed" and data.check_conclusion in ( - "failure", - "cancelled", - "timed_out", - ) - - -def is_pr_merged(data: GitHubWebhookData) -> bool: - """Check if the webhook indicates a PR was merged. - - Args: - data: Parsed webhook data. - - Returns: - True if PR was merged. - """ - return ( - data.event_type == "pull_request" - and data.action == "closed" - and data.raw_payload.get("pull_request", {}).get("merged", False) - ) - - -def is_pr_review_approved(data: GitHubWebhookData) -> bool: - """Check if the webhook indicates PR review approval. - - Args: - data: Parsed webhook data. - - Returns: - True if PR was approved. - """ - return ( - data.event_type == "pull_request_review" - and data.action == "submitted" - and data.raw_payload.get("review", {}).get("state") == "approved" - ) - - -def is_pr_review_changes_requested(data: GitHubWebhookData) -> bool: - """Check if the webhook indicates changes were requested. - - Args: - data: Parsed webhook data. - - Returns: - True if changes were requested. - """ - return ( - data.event_type == "pull_request_review" - and data.action == "submitted" - and data.raw_payload.get("review", {}).get("state") in ("changes_requested", "commented") - ) diff --git a/src/forge/integrations/source_control/__init__.py b/src/forge/integrations/source_control/__init__.py new file mode 100644 index 000000000..01f2842e2 --- /dev/null +++ b/src/forge/integrations/source_control/__init__.py @@ -0,0 +1,71 @@ +"""Provider-neutral source control contracts, errors, and registry.""" + +from forge.integrations.source_control.contracts import ( + Actor, + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + CheckConclusion, + CheckRun, + CheckStatus, + Connection, + EventKind, + NormalizedEvent, + Provider, + RepositoryRef, + RepositoryResolver, + ResolvedRepository, + Review, + ReviewComment, + ReviewState, + SourceControlProvider, + WriteTarget, +) +from forge.integrations.source_control.errors import ( + AuthenticationError, + ConflictError, + NotFoundError, + ProviderConfigError, + RateLimitedError, + SourceControlError, + TransientProviderError, +) +from forge.integrations.source_control.registry import ( + Registry, + get_registry, + load_registry, + register_adapter_factory, +) + +__all__ = [ + "Actor", + "AuthenticationError", + "ChangeRequest", + "ChangeRequestIdentity", + "ChangeRequestState", + "CheckConclusion", + "CheckRun", + "CheckStatus", + "ConflictError", + "Connection", + "EventKind", + "NormalizedEvent", + "NotFoundError", + "Provider", + "ProviderConfigError", + "RateLimitedError", + "Registry", + "RepositoryRef", + "RepositoryResolver", + "ResolvedRepository", + "Review", + "ReviewComment", + "ReviewState", + "SourceControlError", + "SourceControlProvider", + "TransientProviderError", + "WriteTarget", + "get_registry", + "load_registry", + "register_adapter_factory", +] diff --git a/src/forge/integrations/source_control/contracts.py b/src/forge/integrations/source_control/contracts.py new file mode 100644 index 000000000..5713034e6 --- /dev/null +++ b/src/forge/integrations/source_control/contracts.py @@ -0,0 +1,337 @@ +"""Provider-neutral source control contracts. + +Normalized data models and the SourceControlProvider protocol +that every adapter (GitHub, later GitLab) implements. This module has no +I/O and no provider-specific imports. +""" + +from dataclasses import dataclass, field +from datetime import datetime +from enum import StrEnum +from typing import Any, Literal, Protocol, runtime_checkable + + +class Provider(StrEnum): + """Source control providers Forge knows the name of.""" + + GITHUB = "github" + GITLAB = "gitlab" + + +class ChangeRequestState(StrEnum): + OPEN = "open" + CLOSED = "closed" + MERGED = "merged" + + +class ReviewState(StrEnum): + APPROVED = "approved" + CHANGES_REQUESTED = "changes_requested" + COMMENTED = "commented" + PENDING = "pending" + DISMISSED = "dismissed" + + +class CheckStatus(StrEnum): + QUEUED = "queued" + IN_PROGRESS = "in_progress" + COMPLETED = "completed" + + +class CheckConclusion(StrEnum): + SUCCESS = "success" + FAILURE = "failure" + CANCELLED = "cancelled" + SKIPPED = "skipped" + NEUTRAL = "neutral" + NONE = "none" + + +class EventKind(StrEnum): + CR_OPENED = "cr_opened" + CR_UPDATED = "cr_updated" + CR_CLOSED = "cr_closed" + CR_MERGED = "cr_merged" + CHECK_UPDATED = "check_updated" + COMMENT_CREATED = "comment_created" + REVIEW_SUBMITTED = "review_submitted" + PUSH = "push" + UNKNOWN = "unknown" + + +@dataclass(frozen=True) +class Actor: + login: str + is_bot: bool + + +@dataclass(frozen=True) +class RepositoryRef: + id: str + provider: Provider + connection: str + namespace: str + default_branch: str + change_request_mode: Literal["fork", "direct"] + + +@dataclass(frozen=True) +class Connection: + name: str + provider: Provider + base_url: str + credential_env: str + webhook_secret_env: str + ca_path: str | None = None + allowed_namespaces: list[str] | None = None + + +@dataclass(frozen=True) +class ChangeRequestIdentity: + connection: str + repository_id: str + native_id: str | int | None = None + + +@dataclass +class ChangeRequest: + identity: ChangeRequestIdentity + url: str + title: str + body: str + state: ChangeRequestState + source_branch: str + target_branch: str + head_sha: str = "" + draft: bool = False + # Whether create_change_request just created this change request, as opposed + # to returning a pre-existing one for the same head/base pair. Meaningless + # (and always True) for get/update_change_request results. + created: bool = True + + +@dataclass +class ReviewComment: + id: str + body: str + author: str + path: str | None = None + line: int | None = None + resolved: bool = False + in_reply_to: str | None = None + + +@dataclass +class Review: + id: str + state: ReviewState + body: str + author: str + comments: list[ReviewComment] = field(default_factory=list) + + +@dataclass +class CheckRun: + name: str + status: CheckStatus + conclusion: CheckConclusion + url: str | None = None + logs_url: str | None = None + output: dict[str, str] = field(default_factory=dict) # title/summary/text + + +@dataclass(frozen=True) +class GitCredentials: + """Everything local `git` (clone/fetch/push over HTTPS) needs to talk to + a connection's host, independent of that connection's API surface. + + Returned by ``SourceControlProvider.get_git_credentials`` -- a cheap, + side-effect-free derivation from the adapter's already-resolved + Connection/credential, safe to call from any code path (including one + reconstructing a workspace from persisted state, not a fresh API call). + """ + + host: str # bare host, e.g. "github.com" or "ghe.example.com" -- no scheme + # repr=False: keep the raw token out of the dataclass's default repr, so + # an unhandled exception's traceback or a stray `logger.debug(credentials)` + # doesn't print it -- this module deliberately avoids a pydantic/SecretStr + # dependency (see the module docstring), so a repr guard is the + # lightweight equivalent. + token: str = field(repr=False) + # CA bundle for a connection's self-signed TLS certificate (GitHub + # Enterprise Server). None for the common case (public GitHub or a CA + # trusted by the default store). + ca_path: str | None = None + + +@dataclass +class WriteTarget: + clone_url: str + push_remote_name: str + head_ref: str # the source branch to open the change request from + base_branch: str + # Fork identity, populated only for change_request_mode == "fork"; None for + # "direct". Callers that push via local git need these to add the fork remote + # and build the provider-native cross-fork head ref. + fork_owner: str | None = None + fork_repo: str | None = None + + +@dataclass +class NormalizedEvent: + id: str + kind: EventKind + repo_ref: RepositoryRef + actor: Actor + received_at: datetime + change_request: ChangeRequest | None = None + comment: ReviewComment | None = None + review: Review | None = None + check: CheckRun | None = None + # Status of the check *suite* a CHECK_UPDATED event belongs to, independent + # of `check` (which is only populated for individual check_run events, not + # check_suite events). None when the event has no associated suite status. + check_suite_status: CheckStatus | None = None + raw: dict[str, Any] = field(default_factory=dict) + + +@runtime_checkable +class SourceControlProvider(Protocol): + """Provider-neutral operations every source control adapter implements.""" + + async def verify_webhook(self, headers: dict[str, str], body: bytes) -> bool: ... + + async def parse_webhook( + self, headers: dict[str, str], body: bytes, resolver: "RepositoryResolver" + ) -> NormalizedEvent: ... + + async def resolve_default_branch(self, repo_ref: RepositoryRef) -> str: ... + + async def get_git_credentials(self, repo_ref: RepositoryRef) -> GitCredentials: + """Host/token/CA for local `git` operations against this connection. + + Side-effect-free (no API call) -- a pure derivation from the + adapter's already-resolved Connection/credential. Declared async + (despite never awaiting) to match the rest of this protocol: a sync + method here is a footgun against AsyncMock-based test doubles, which + silently hand back an unawaited coroutine instead of raising. Safe to + call from any code path, including one reconstructing GitOperations + from persisted workflow state rather than a fresh repo_ref + resolution. + """ + ... + + async def ensure_write_target(self, repo_ref: RepositoryRef) -> WriteTarget: ... + + async def create_change_request( + self, + repo_ref: RepositoryRef, + target: WriteTarget, + title: str, + body: str, + draft: bool = False, + ) -> ChangeRequest: ... + + async def get_change_request( + self, repo_ref: RepositoryRef, identity: ChangeRequestIdentity + ) -> ChangeRequest: ... + + async def update_change_request( + self, + repo_ref: RepositoryRef, + identity: ChangeRequestIdentity, + *, + title: str | None = None, + body: str | None = None, + state: ChangeRequestState | None = None, + ) -> ChangeRequest: ... + + async def create_comment( + self, repo_ref: RepositoryRef, identity: ChangeRequestIdentity, body: str + ) -> ReviewComment: ... + + async def reply_to_comment( + self, + repo_ref: RepositoryRef, + identity: ChangeRequestIdentity, + comment_id: str, + body: str, + ) -> ReviewComment: ... + + async def get_review_threads( + self, repo_ref: RepositoryRef, identity: ChangeRequestIdentity + ) -> list[Review]: + """Submission-level review verdicts (approve/request-changes/comment); comments empty.""" + ... + + async def get_review_thread_comments( + self, repo_ref: RepositoryRef, identity: ChangeRequestIdentity + ) -> list[Review]: + """Unresolved inline diff-comment threads; one Review per thread, comments populated.""" + ... + + async def get_review_comments_for_submission( + self, repo_ref: RepositoryRef, identity: ChangeRequestIdentity, review_id: str + ) -> list[ReviewComment]: + """Inline comments from one specific review submission. + + Scoped to a single review, unlike get_review_thread_comments (every + unresolved thread regardless of which review raised it) -- avoids + pulling in stale comments from prior review rounds on the same PR. + """ + ... + + async def get_checks(self, repo_ref: RepositoryRef, ref: str) -> list[CheckRun]: ... + + async def get_check_logs(self, repo_ref: RepositoryRef, check: CheckRun) -> str: ... + + async def get_check_artifacts( + self, repo_ref: RepositoryRef, check: CheckRun + ) -> list[tuple[str, bytes]]: ... + + async def get_file(self, repo_ref: RepositoryRef, path: str, ref: str) -> str: ... + + async def put_file( + self, + repo_ref: RepositoryRef, + path: str, + content: str, + message: str, + branch: str, + ) -> None: ... + + async def create_branch(self, repo_ref: RepositoryRef, name: str, base: str) -> None: ... + + async def get_authenticated_identity(self, repo_ref: RepositoryRef) -> Actor: ... + + async def close(self) -> None: + """Release this adapter's underlying HTTP client/connection pool. + + Called once by Registry.aclose() at process shutdown for every + adapter it has cached -- not per-operation. A no-op if the adapter + never lazily constructed a client (i.e. was never actually used). + """ + ... + + +@dataclass(frozen=True) +class ResolvedRepository: + """Resolving an identifier yields a repository, its connection, and (if a provider has + registered one) an adapter instance. `adapter` is None until a provider registers a + factory (see registry.register_adapter_factory, added in Task 5).""" + + repo_ref: RepositoryRef + connection: Connection + adapter: SourceControlProvider | None = None + + +class RepositoryResolver(Protocol): + """Structural interface that registry.Registry satisfies. + + Declared here (not imported from registry.py) so SourceControlProvider.parse_webhook + can reference it without contracts.py importing registry.py. + """ + + def resolve( + self, identifier: str, provider_hint: Provider | None = None + ) -> ResolvedRepository: ... diff --git a/src/forge/integrations/source_control/errors.py b/src/forge/integrations/source_control/errors.py new file mode 100644 index 000000000..94b3a0965 --- /dev/null +++ b/src/forge/integrations/source_control/errors.py @@ -0,0 +1,38 @@ +"""Provider-neutral source control exception hierarchy. + +Adapters translate provider-specific failures (HTTP status codes, GraphQL +errors) into these at the boundary; nothing above the adapter layer touches +a provider's own exception types directly. +""" + + +class SourceControlError(Exception): + """Base exception for all source control provider errors.""" + + +class AuthenticationError(SourceControlError): + """Raised when a provider rejects the configured credentials.""" + + +class NotFoundError(SourceControlError): + """Raised when a repository, connection, or provider resource can't be resolved.""" + + +class ConflictError(SourceControlError): + """Raised for PR-already-exists or fork-diverged-on-sync conditions.""" + + +class RateLimitedError(SourceControlError): + """Raised when a provider rate-limits a request.""" + + def __init__(self, message: str, retry_after: float | None = None) -> None: + super().__init__(message) + self.retry_after = retry_after + + +class TransientProviderError(SourceControlError): + """Raised for retryable 5xx/timeout provider failures.""" + + +class ProviderConfigError(SourceControlError): + """Raised for invalid connection/repository configuration.""" diff --git a/src/forge/integrations/source_control/github/__init__.py b/src/forge/integrations/source_control/github/__init__.py new file mode 100644 index 000000000..b2ac85a80 --- /dev/null +++ b/src/forge/integrations/source_control/github/__init__.py @@ -0,0 +1,32 @@ +"""GitHub source control adapter.""" + +from forge.config import get_settings +from forge.integrations.source_control.contracts import Connection, Provider +from forge.integrations.source_control.github.adapter import GitHubAdapter +from forge.integrations.source_control.registry import ( + register_adapter_factory, + resolve_env_value, +) + +__all__ = ["GitHubAdapter"] + + +def _build_github_adapter(connection: Connection) -> GitHubAdapter: + """Registry factory: resolve the connection's credential/secret and bind them. + + The connection stores only the *names* of the env vars (spec 154). Resolve + them through Settings-aware lookup so a .env-only GITHUB_TOKEN is honored. + """ + settings = get_settings() + credential = resolve_env_value(connection.credential_env, settings) + webhook_secret = ( + resolve_env_value(connection.webhook_secret_env, settings) + if connection.webhook_secret_env + else None + ) + return GitHubAdapter( + connection=connection, credential=credential, webhook_secret=webhook_secret + ) + + +register_adapter_factory(Provider.GITHUB, _build_github_adapter) diff --git a/src/forge/integrations/source_control/github/adapter.py b/src/forge/integrations/source_control/github/adapter.py new file mode 100644 index 000000000..980db1dab --- /dev/null +++ b/src/forge/integrations/source_control/github/adapter.py @@ -0,0 +1,1243 @@ +"""GitHub implementation of the SourceControlProvider protocol.""" + +import base64 +import functools +import hashlib +import hmac +import json +import logging +from datetime import UTC, datetime + +import httpx +from pydantic import SecretStr + +from forge.config import get_settings +from forge.integrations.github.client import GitHubClient +from forge.integrations.source_control.contracts import ( + Actor, + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + CheckConclusion, + CheckRun, + CheckStatus, + Connection, + EventKind, + GitCredentials, + NormalizedEvent, + Provider, + RepositoryRef, + RepositoryResolver, + Review, + ReviewComment, + ReviewState, + WriteTarget, +) +from forge.integrations.source_control.errors import ( + AuthenticationError, + ConflictError, + NotFoundError, + RateLimitedError, + SourceControlError, + TransientProviderError, +) + +logger = logging.getLogger(__name__) + +_REVIEW_STATE_MAP: dict[str, ReviewState] = { + "APPROVED": ReviewState.APPROVED, + "CHANGES_REQUESTED": ReviewState.CHANGES_REQUESTED, + "COMMENTED": ReviewState.COMMENTED, + "PENDING": ReviewState.PENDING, + "DISMISSED": ReviewState.DISMISSED, +} + +_CHECK_STATUS_MAP: dict[str, CheckStatus] = { + "queued": CheckStatus.QUEUED, + "in_progress": CheckStatus.IN_PROGRESS, + "completed": CheckStatus.COMPLETED, +} + +_CHECK_CONCLUSION_MAP: dict[str, CheckConclusion] = { + "success": CheckConclusion.SUCCESS, + "failure": CheckConclusion.FAILURE, + "timed_out": CheckConclusion.FAILURE, + "action_required": CheckConclusion.FAILURE, + "cancelled": CheckConclusion.CANCELLED, + "skipped": CheckConclusion.SKIPPED, + "neutral": CheckConclusion.NEUTRAL, +} + +_GITHUB_ACTIONS_APP_SLUG = "github-actions" +_DEFAULT_API_BASE_URL = "https://api.github.com" +_DEFAULT_WEB_BASE_URL = "https://github.com" + + +def _web_base_url(connection: Connection) -> str: + """Derive the git/web host from a connection's API base_url. + + Public GitHub's API and web hosts differ (api.github.com vs github.com). + GitHub Enterprise Server's API root is the same host with an /api/v3 + suffix, so its web root is that host with the suffix stripped. + """ + base = (connection.base_url or "").rstrip("/") + if not base or base == _DEFAULT_API_BASE_URL: + return _DEFAULT_WEB_BASE_URL + if base.endswith("/api/v3"): + return base[: -len("/api/v3")] + return base + + +def _looks_like_stale_sha_error(response: httpx.Response) -> bool: + """Best-effort check that a 422 from the Contents API is a stale-sha + conflict, not an unrelated validation failure (branch protection, the + path being a directory, etc.) that ConflictError's "retry with a fresh + read" guidance would misrepresent. + """ + try: + message = str(response.json().get("message", "")).lower() + except Exception: + return False + return "sha" in message + + +def _translate_provider_errors(func): + """Translate a raw ``httpx.HTTPStatusError`` into the neutral error hierarchy. + + Applied to every adapter method that calls the GitHub API, so the boundary + contract documented on errors.py -- "nothing above the adapter layer + touches a provider's own exception types directly" -- actually holds + everywhere, not just in the couple of methods that used to hand-roll this. + Methods that already translate specific statuses themselves (get_check_logs, + put_file) re-raise anything they don't recognize, which this then maps too. + Statuses with no generic neutral mapping (404, 409, 422, ...) are left for + callers to handle themselves and propagate unchanged. + """ + + @functools.wraps(func) + async def wrapper(*args, **kwargs): + try: + return await func(*args, **kwargs) + except httpx.HTTPStatusError as exc: + status = exc.response.status_code + if status in (401, 403): + raise AuthenticationError(f"GitHub rejected the request: {exc}") from exc + if status == 429: + retry_after = exc.response.headers.get("Retry-After") + try: + parsed_retry_after = float(retry_after) if retry_after else None + except ValueError: + parsed_retry_after = None + raise RateLimitedError( + f"GitHub rate-limited the request: {exc}", + retry_after=parsed_retry_after, + ) from exc + if status >= 500: + raise TransientProviderError(f"GitHub returned {status}: {exc}") from exc + raise + except httpx.TransportError as exc: + # Covers httpx.TimeoutException, ConnectError, and other network-level + # failures -- all subclasses of TransportError. + raise TransientProviderError(f"GitHub request failed: {exc}") from exc + + return wrapper + + +def _require_native_id(identity: ChangeRequestIdentity) -> int: + """Coerce identity.native_id to an int, rejecting a missing/None value. + + A None native_id (e.g. from an identity built off a malformed webhook + payload) must not silently become PR number 0 and produce a confusing + 404 from GitHub — fail clearly at the adapter boundary instead. + """ + if identity.native_id is None: + raise ValueError(f"ChangeRequestIdentity has no native_id: {identity}") + return int(identity.native_id) + + +class GitHubAdapter: + """GitHub implementation of SourceControlProvider protocol.""" + + def __init__( + self, + connection: Connection, + credential: str | None = None, + webhook_secret: str | None = None, + client: GitHubClient | None = None, + ): + """Initialize GitHub adapter. + + Args: + connection: Connection configuration + credential: Already-resolved GitHub token to use for API calls. This + constructor does not itself read the environment variable named by + connection.credential_env; the caller is responsible for resolving + that value and passing it in here. If None, the lazily-constructed + client falls back to the process-wide Settings/GITHUB_TOKEN. + webhook_secret: Already-resolved webhook secret used to verify inbound + webhooks. This constructor does not itself read the environment + variable named by connection.webhook_secret_env; the caller is + responsible for resolving that value and passing it in here. + client: Pre-built GitHubClient to use instead of lazily constructing one. + Primarily for tests that need to inject a mocked httpx.AsyncClient. + """ + self._connection = connection + self._credential = credential + self._webhook_secret = webhook_secret + self._client: GitHubClient | None = client + + def _get_client(self) -> GitHubClient: + """Get the injected GitHub client, or lazily construct one. + + When no client was injected, a new GitHubClient is built using this + adapter's configured credential (if any) so that adapter instances + configured for different connections/credentials don't all collapse + onto the process-wide Settings/GITHUB_TOKEN. + """ + if self._client is None: + if self._credential is not None: + settings = get_settings().model_copy( + update={"github_token": SecretStr(self._credential)} + ) + self._client = GitHubClient( + settings=settings, + base_url=self._connection.base_url or None, + ca_path=self._connection.ca_path, + ) + else: + self._client = GitHubClient( + base_url=self._connection.base_url or None, + ca_path=self._connection.ca_path, + ) + return self._client + + async def verify_webhook(self, headers: dict[str, str], body: bytes) -> bool: + """Verify GitHub webhook signature. + + Args: + headers: Request headers. X-Hub-Signature-256 is required when a + webhook secret is configured. + body: Raw request body bytes + + Returns: + True if signature is valid, False otherwise + """ + if not self._webhook_secret: + logger.warning( + "GitHub webhook signature verification disabled: no webhook secret " + "configured for this connection" + ) + return True + + signature_header = headers.get("X-Hub-Signature-256", "") + if not signature_header.startswith("sha256="): + return False + + expected_signature = signature_header[len("sha256=") :] + + computed = hmac.new(self._webhook_secret.encode(), body, hashlib.sha256).hexdigest() + + return hmac.compare_digest(computed, expected_signature) + + async def parse_webhook( + self, + headers: dict[str, str], + body: bytes, + resolver: RepositoryResolver, + ) -> NormalizedEvent: + """Parse GitHub webhook into NormalizedEvent. + + Args: + headers: Request headers (must include X-GitHub-Event, X-GitHub-Delivery) + body: Raw request body bytes + resolver: Registry for resolving repository references + + Returns: + Normalized event + """ + event_type = headers.get("X-GitHub-Event", "") + event_id = headers.get("X-GitHub-Delivery", "") + + payload = json.loads(body.decode()) + + repo_namespace = payload.get("repository", {}).get("full_name", "") + + resolved = resolver.resolve(repo_namespace, provider_hint=Provider.GITHUB) + repo_ref = resolved.repo_ref + + sender = payload.get("sender", {}) + sender_login = sender.get("login", "") + actor = Actor( + login=sender_login, + is_bot=sender.get("type") == "Bot" or "[bot]" in sender_login, + ) + + kind = self._map_event_kind(event_type, payload) + + change_request = None + pr_payload = payload.get("pull_request") + issue_pr_stub = payload.get("issue", {}).get("pull_request") + check_pr_stub = self._extract_check_event_pull_request(payload, event_type) + if pr_payload is not None: + change_request = self._map_change_request(pr_payload, repo_ref=repo_ref) + elif issue_pr_stub is not None: + issue = payload["issue"] + # The `issue` payload on an issue_comment event has no `merged` field, + # so a comment on an already-merged PR is reported as CLOSED rather + # than MERGED here. Callers that need to distinguish the two must + # fetch the PR directly. + change_request = ChangeRequest( + identity=ChangeRequestIdentity( + connection=repo_ref.connection, + repository_id=repo_ref.id, + native_id=issue.get("number"), + ), + url=issue.get("html_url", ""), + title=issue.get("title", ""), + body=issue.get("body", "") or "", + state=self._map_pr_state(issue.get("state"), False), + source_branch="", + target_branch="", + head_sha="", + draft=False, + ) + elif check_pr_stub is not None: + change_request = ChangeRequest( + identity=ChangeRequestIdentity( + connection=repo_ref.connection, + repository_id=repo_ref.id, + native_id=check_pr_stub.get("number"), + ), + url="", + title="", + body="", + state=ChangeRequestState.OPEN, + source_branch=check_pr_stub.get("head", {}).get("ref", ""), + target_branch=check_pr_stub.get("base", {}).get("ref", ""), + head_sha=check_pr_stub.get("head", {}).get("sha", ""), + draft=False, + ) + + comment = None + if event_type in ("issue_comment", "pull_request_review_comment") and "comment" in payload: + raw_comment = payload["comment"] + in_reply_to = raw_comment.get("in_reply_to_id") + comment = self._map_review_comment( + raw_comment, in_reply_to=str(in_reply_to) if in_reply_to is not None else None + ) + + review = None + if event_type == "pull_request_review" and "review" in payload: + review = self._map_review(payload["review"]) + + check = None + if event_type == "check_run" and "check_run" in payload: + check = self._map_check_run(payload["check_run"]) + + check_suite_status = self._extract_check_suite_status(payload, event_type) + + return NormalizedEvent( + id=event_id, + kind=kind, + repo_ref=repo_ref, + actor=actor, + received_at=datetime.now(UTC), + change_request=change_request, + comment=comment, + review=review, + check=check, + check_suite_status=check_suite_status, + raw=payload, + ) + + def _extract_check_suite_status(self, payload: dict, event_type: str) -> CheckStatus | None: + """Status of the check_suite a check_run/check_suite webhook belongs to. + + check_suite events carry the status on the suite itself; check_run + events carry it on the run's nested check_suite. Used by worker.py to + gate CI-cycle completeness without reading the raw GitHub payload + directly. + """ + if event_type == "check_suite": + raw_status = payload.get("check_suite", {}).get("status") + elif event_type == "check_run": + raw_status = payload.get("check_run", {}).get("check_suite", {}).get("status") + else: + raw_status = None + return _CHECK_STATUS_MAP.get(raw_status) if raw_status else None + + def _map_event_kind(self, event_type: str, payload: dict) -> EventKind: + """Map GitHub event type to normalized EventKind.""" + action = payload.get("action", "") + + if event_type == "pull_request": + if action == "opened": + return EventKind.CR_OPENED + if action in ("synchronize", "edited", "reopened"): + return EventKind.CR_UPDATED + if action == "closed": + if payload.get("pull_request", {}).get("merged", False): + return EventKind.CR_MERGED + return EventKind.CR_CLOSED + + if event_type in ("check_run", "check_suite"): + return EventKind.CHECK_UPDATED + + if event_type in ("issue_comment", "pull_request_review_comment"): + if action == "created": + return EventKind.COMMENT_CREATED + return EventKind.UNKNOWN + + if event_type == "pull_request_review": + if action == "submitted": + return EventKind.REVIEW_SUBMITTED + return EventKind.UNKNOWN + + if event_type == "push": + return EventKind.PUSH + + return EventKind.UNKNOWN + + def _map_pr_state(self, state: str | None, merged: bool) -> ChangeRequestState: + """Map GitHub PR state to normalized ChangeRequestState.""" + if merged: + return ChangeRequestState.MERGED + if state == "closed": + return ChangeRequestState.CLOSED + return ChangeRequestState.OPEN + + def _extract_check_event_pull_request(self, payload: dict, event_type: str) -> dict | None: + """Find the "simple pull request" stub a check_run/check_suite webhook + carries, if any. check_suite events list PRs on the suite itself; + check_run events list them on the run, falling back to the run's + nested check_suite (both are real GitHub payload shapes).""" + if event_type == "check_suite": + pull_requests = payload.get("check_suite", {}).get("pull_requests", []) + elif event_type == "check_run": + check_run = payload.get("check_run", {}) + pull_requests = check_run.get("pull_requests") or check_run.get("check_suite", {}).get( + "pull_requests", [] + ) + else: + pull_requests = [] + if len(pull_requests) > 1: + # Multiple open PRs can share a head branch/SHA (stacked PRs, + # re-targeted base branch). We have no reliable way to disambiguate + # here, so the first entry is used and the rest are silently + # ignored -- log so a misattributed CI status is traceable. + logger.warning( + "check event lists %d pull_requests; using the first (#%s) " + "and ignoring the rest: %s", + len(pull_requests), + pull_requests[0].get("number"), + [pr.get("number") for pr in pull_requests[1:]], + ) + return pull_requests[0] if pull_requests else None + + @_translate_provider_errors + async def resolve_default_branch(self, repo_ref: RepositoryRef) -> str: + """Get the default branch name for a repository. + + Args: + repo_ref: Repository reference + + Returns: + Default branch name (e.g. "main") + """ + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + repo_data = await client.get_repository(owner, repo) + return repo_data.get("default_branch", "main") + + async def get_git_credentials(self, _repo_ref: RepositoryRef) -> GitCredentials: + """Host/token/CA for local `git` operations against this connection. + + Pure derivation from this adapter's already-resolved Connection and + credential -- no API call, so it's safe to call from any code path, + including one reconstructing GitOperations from persisted state. + + Args: + _repo_ref: Repository reference (unused -- credentials are + per-connection, not per-repo; part of the protocol for + providers that might scope credentials more narrowly). + """ + web_base = _web_base_url(self._connection) + host = web_base.removeprefix("https://").removeprefix("http://") + token = self._credential or get_settings().github_token.get_secret_value() + return GitCredentials(host=host, token=token, ca_path=self._connection.ca_path) + + @_translate_provider_errors + async def ensure_write_target(self, repo_ref: RepositoryRef) -> WriteTarget: + """Ensure a place to push commits exists before opening a change request. + + For ``change_request_mode == "direct"`` the upstream repository itself is + the write target and nothing needs to be created. For the default + ``"fork"`` mode, a fork is created (or reused) and synced with upstream so + the fork's default branch matches before Forge branches off it. + + Args: + repo_ref: Repository reference. + + Returns: + WriteTarget describing where to push and which branch to base off. + + Raises: + ConflictError: In fork mode, when the fork has diverged from upstream + and GitHub's merge-upstream API cannot fast-forward it. + """ + owner, repo = repo_ref.namespace.split("/", 1) + web_base = _web_base_url(self._connection) + + if repo_ref.change_request_mode == "direct": + return WriteTarget( + clone_url=f"{web_base}/{owner}/{repo}.git", + push_remote_name="origin", + head_ref=f"forge/{owner}/{repo}", + base_branch=repo_ref.default_branch, + ) + + client = self._get_client() + + fork = await client.get_or_create_fork(owner, repo) + fork_owner = fork["owner"]["login"] + fork_name = fork["name"] + + synced = await client.sync_fork_with_upstream( + fork_owner, + fork_name, + branch=repo_ref.default_branch, + ) + if not synced: + raise ConflictError( + f"Fork {fork_owner}/{fork_name}:{repo_ref.default_branch} has " + f"diverged from upstream {owner}/{repo} and cannot be auto-synced." + ) + + clone_url = fork.get("clone_url") or f"{web_base}/{fork_owner}/{fork_name}.git" + return WriteTarget( + clone_url=clone_url, + push_remote_name="origin", + head_ref=f"forge/{owner}/{repo}", + base_branch=repo_ref.default_branch, + fork_owner=fork_owner, + fork_repo=fork_name, + ) + + @_translate_provider_errors + async def create_change_request( + self, + repo_ref: RepositoryRef, + target: WriteTarget, + title: str, + body: str, + draft: bool = False, + ) -> ChangeRequest: + """Create a pull request (or reuse an existing one for the same branches). + + Args: + repo_ref: Repository reference. + target: Write target produced by ensure_write_target; supplies the + head/base branches to open the PR against. + title: PR title. + body: PR description. + draft: Whether to open the PR as a draft. + + Returns: + The created (or already-existing) change request. When GitHub reports + a PR already exists for the head/base pair, the client returns the + existing PR (created=False) and it is mapped and returned as-is. + """ + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + + head = f"{target.fork_owner}:{target.head_ref}" if target.fork_owner else target.head_ref + result = await client.create_pull_request( + owner=owner, + repo=repo, + title=title, + body=body, + head=head, + base=target.base_branch, + draft=draft, + ) + return self._map_change_request(result.pr, repo_ref=repo_ref, created=result.created) + + @_translate_provider_errors + async def get_change_request( + self, repo_ref: RepositoryRef, identity: ChangeRequestIdentity + ) -> ChangeRequest: + """Fetch a pull request by identity. + + Args: + repo_ref: Repository reference. + identity: Change request identity carrying the PR number in native_id. + + Returns: + The change request, preserving the identity passed in. + """ + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + + pr = await client.get_pull_request(owner, repo, _require_native_id(identity)) + return self._map_change_request(pr, identity=identity) + + @_translate_provider_errors + async def update_change_request( + self, + repo_ref: RepositoryRef, + identity: ChangeRequestIdentity, + *, + title: str | None = None, + body: str | None = None, + state: ChangeRequestState | None = None, + ) -> ChangeRequest: + """Update a pull request's title, body, and/or state. + + Args: + repo_ref: Repository reference. + identity: Change request identity carrying the PR number in native_id. + title: New title, if changing. + body: New body, if changing. + state: New state, if changing. Only OPEN and CLOSED are settable via + GitHub's PR update endpoint. + + Returns: + The updated change request, preserving the identity passed in. + + Raises: + ValueError: If state is ChangeRequestState.MERGED. GitHub's PR update + (PATCH) endpoint cannot set "merged"; merging is a distinct + operation via the merge endpoint (not modelled here). Raising is + preferred over silently ignoring the request or silently closing + the PR instead, both of which would misrepresent what happened. + """ + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + + gh_state: str | None = None + if state is not None: + if state == ChangeRequestState.MERGED: + raise ValueError( + "Cannot set change request state to MERGED via update_change_request; " + "GitHub's PR update endpoint only supports 'open' or 'closed'. " + "Merging requires the separate merge operation." + ) + gh_state = state.value + + pr = await client.update_pull_request( + owner=owner, + repo=repo, + pr_number=_require_native_id(identity), + title=title, + body=body, + state=gh_state, + ) + return self._map_change_request(pr, identity=identity) + + def _map_change_request( + self, + pr: dict, + *, + repo_ref: RepositoryRef | None = None, + identity: ChangeRequestIdentity | None = None, + created: bool = False, + ) -> ChangeRequest: + """Map a GitHub PR dict into a ChangeRequest. + + Exactly one of ``repo_ref`` (to construct a fresh identity from the PR + number) or ``identity`` (to preserve a caller-supplied identity) must be + given. ``created`` is only meaningful from create_change_request, which + passes through whether GitHub actually opened a new PR or returned a + pre-existing one for the same head/base pair. + """ + if repo_ref is not None and identity is not None: + raise ValueError("_map_change_request accepts repo_ref or identity, not both") + if identity is None: + if repo_ref is None: + raise ValueError("_map_change_request requires either repo_ref or identity") + identity = ChangeRequestIdentity( + connection=repo_ref.connection, + repository_id=repo_ref.id, + native_id=pr.get("number"), + ) + + return ChangeRequest( + identity=identity, + url=pr.get("html_url", ""), + title=pr.get("title", ""), + body=pr.get("body", "") or "", + state=self._map_pr_state(pr.get("state"), pr.get("merged", False)), + source_branch=pr.get("head", {}).get("ref", ""), + target_branch=pr.get("base", {}).get("ref", ""), + head_sha=pr.get("head", {}).get("sha", ""), + draft=pr.get("draft", False), + created=created, + ) + + @_translate_provider_errors + async def create_comment( + self, repo_ref: RepositoryRef, identity: ChangeRequestIdentity, body: str + ) -> ReviewComment: + """Create a general (PR-level) comment on a change request. + + The protocol's create_comment carries no path/line/commit_id, so this + can only post a general issue-style comment on the PR conversation, not + an inline diff comment. The client prepends the bot prefix internally. + + Args: + repo_ref: Repository reference. + identity: Change request identity carrying the PR number in native_id. + body: Comment text. + + Returns: + The created comment mapped to a ReviewComment. + """ + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + + comment = await client.create_issue_comment( + owner=owner, + repo=repo, + issue_number=_require_native_id(identity), + body=body, + ) + return self._map_review_comment(comment) + + @_translate_provider_errors + async def reply_to_comment( + self, + repo_ref: RepositoryRef, + identity: ChangeRequestIdentity, + comment_id: str, + body: str, + ) -> ReviewComment: + """Reply within the review thread containing an existing review comment. + + Uses GitHub's review-comment reply endpoint so the reply lands in the + actual thread rather than as a detached PR-level comment. The client + prepends the bot prefix internally. + + Args: + repo_ref: Repository reference. + identity: Change request identity carrying the PR number in native_id. + comment_id: ID of the review comment to reply to. + body: Reply text. + + Returns: + The created reply mapped to a ReviewComment, with in_reply_to set to + the parent comment_id. + """ + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + + comment = await client.reply_to_review_comment( + owner=owner, + repo=repo, + pr_number=_require_native_id(identity), + comment_id=int(comment_id), + body=body, + ) + return self._map_review_comment(comment, in_reply_to=comment_id) + + def _map_review_comment( + self, comment: dict, *, in_reply_to: str | None = None + ) -> ReviewComment: + """Map a GitHub comment API response into a ReviewComment.""" + return ReviewComment( + id=str(comment.get("id", "")), + body=comment.get("body", "") or "", + # GitHub returns "user": null for comments from a deleted account, so a + # present-but-null key must not be treated as "absent" (.get's default + # only applies when the key itself is missing). + author=(comment.get("user") or {}).get("login", ""), + path=comment.get("path"), + line=comment.get("line"), + in_reply_to=in_reply_to, + ) + + @_translate_provider_errors + async def get_review_threads( + self, repo_ref: RepositoryRef, identity: ChangeRequestIdentity + ) -> list[Review]: + """Get the formal review submissions on a pull request. + + Despite the protocol name, this returns GitHub's review-level verdicts + (approve / request-changes / comment), each mapped to a Review with its + submission-level ReviewState. + + Each Review's ``comments`` is intentionally left empty. GitHub's inline + review comments (fetched via get_pull_request_review_threads / + get_pull_request_review_comments) don't carry a reliable link back to + which review submission they belong to in the data GitHubClient + currently exposes — the REST comments payload has a + ``pull_request_review_id`` field, but the existing thread-reconciliation + methods don't preserve it. Attaching comments to the correct review + isn't possible without deeper client changes that are out of scope here. + + Args: + repo_ref: Repository reference. + identity: Change request identity carrying the PR number in native_id. + + Returns: + List of Review objects (one per submitted review), with empty comments. + """ + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + + reviews_data = await client.get_reviews(owner, repo, _require_native_id(identity)) + return [self._map_review(review) for review in reviews_data] + + def _map_review(self, review: dict) -> Review: + """Map a GitHub review dict into a Review. + + Shared by get_review_threads (REST /reviews, uppercase state strings + like "APPROVED") and parse_webhook (pull_request_review webhooks, + lowercase state strings like "approved") -- uppercasing before the + lookup makes both inputs safe; it's a no-op for the already-uppercase + REST case. + """ + gh_state = (review.get("state") or "").upper() + state = _REVIEW_STATE_MAP.get(gh_state, ReviewState.COMMENTED) + return Review( + id=str(review.get("id", "")), + state=state, + body=review.get("body", "") or "", + # GitHub returns "user": null for a deleted account, so a + # present-but-null key must not be treated as absent. + author=(review.get("user") or {}).get("login", ""), + comments=[], + ) + + @_translate_provider_errors + async def get_review_thread_comments( + self, repo_ref: RepositoryRef, identity: ChangeRequestIdentity + ) -> list[Review]: + """Return unresolved inline review threads as Reviews with populated comments. + + One Review per thread: Review.id is the thread id, Review.comments are the + inline diff comments. Wraps the client's GraphQL->REST thread + reconciliation (which already filters resolved/outdated threads). + """ + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + threads = await client.get_pull_request_review_threads( + owner, repo, _require_native_id(identity) + ) + reviews: list[Review] = [] + for thread in threads: + comments = [ + ReviewComment( + id=str(c.get("comment_id") or ""), + body=c.get("body", "") or "", + author=c.get("author", "") or "", + path=thread.get("path"), + line=thread.get("line"), + ) + for c in thread.get("comments", []) + ] + reviews.append( + Review( + id=str(thread.get("thread_id") or ""), + state=ReviewState.COMMENTED, + body="", + author=comments[0].author if comments else "", + comments=comments, + ) + ) + return reviews + + @_translate_provider_errors + async def get_review_comments_for_submission( + self, repo_ref: RepositoryRef, identity: ChangeRequestIdentity, review_id: str + ) -> list[ReviewComment]: + """Inline comments from one specific review submission (REST, not the + thread-reconciliation GraphQL query get_review_thread_comments uses).""" + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + comments = await client.get_review_comments( + owner, repo, _require_native_id(identity), int(review_id) + ) + return [ + ReviewComment( + id=str(c.get("id", "")), + body=c.get("body", "") or "", + # GitHub returns "user": null for comments from a deleted account. + author=(c.get("user") or {}).get("login", ""), + path=c.get("path"), + # GitHub returns "line": null for comments anchored to a diff + # position that's no longer part of the current diff (e.g. an + # outdated review comment); original_line/position still carry + # the location in that case. + line=c.get("line") or c.get("original_line") or c.get("position"), + ) + for c in comments + ] + + @_translate_provider_errors + async def get_checks(self, repo_ref: RepositoryRef, ref: str) -> list[CheckRun]: + """Get all CI check results for a ref. + + The underlying client combines two GitHub CI surfaces into one list: + GitHub Actions check runs and legacy commit statuses (Prow, etc.). Both + are mapped into CheckRun here. + + Only check runs created by the GitHub Actions app (``app.slug == + "github-actions"``) have logs fetchable via the Actions job-logs + endpoint. A check-run id is NOT an Actions job id, though -- they address + different resources -- so what gets stored in ``logs_url`` for those is + the workflow *run* id parsed from ``details_url``. get_check_logs later + lists that run's jobs and matches this check by name to recover the real + Actions job id the logs endpoint needs. Check runs from other GitHub + Apps using the Checks API (CodeQL, Codecov, Vercel, etc.) and commit + statuses (Prow, etc.) get ``logs_url = None`` -- the former aren't + Actions jobs, the latter have no logs endpoint at all -- so + get_check_logs has nothing to fetch for them. + + Status/conclusion map from GitHub's string values. An unrecognized or + missing conclusion becomes CheckConclusion.NONE. An unrecognized or + missing status is inferred from the conclusion: a present conclusion + implies the check finished (CheckStatus.COMPLETED), otherwise it is + treated as still running (CheckStatus.IN_PROGRESS). + + Args: + repo_ref: Repository reference. + ref: Git ref (commit SHA, branch, or tag). + + Returns: + List of CheckRun objects. + """ + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + + entries = await client.get_check_runs(owner=owner, repo=repo, ref=ref) + return [self._map_check_run(entry) for entry in entries] + + def _map_check_run(self, entry: dict) -> CheckRun: + """Map a GitHub check-run/commit-status dict into a CheckRun. + + Shared by get_checks (REST, one call per entry in the combined + check-runs/commit-statuses list) and parse_webhook (a single + check_run webhook payload) -- see get_checks' docstring for the + logs_url/app-slug and status/conclusion inference rules this applies. + """ + gh_conclusion = entry.get("conclusion") or "" + conclusion = _CHECK_CONCLUSION_MAP.get(gh_conclusion, CheckConclusion.NONE) + + gh_status = entry.get("status") or "" + status = _CHECK_STATUS_MAP.get(gh_status) + if status is None: + # Missing/unrecognized status: infer from whether the check has a + # conclusion. A concluded check is done; otherwise assume running. + status = ( + CheckStatus.COMPLETED + if conclusion is not CheckConclusion.NONE + else CheckStatus.IN_PROGRESS + ) + + # Only GitHub Actions check runs have fetchable logs, and the logs + # endpoint keys off the Actions *job* id -- a different resource from + # this check run's own id. What is recoverable here is the workflow + # *run* id (from details_url); get_check_logs uses it to list the run's + # jobs and match this check by name to the real job id. Store the run id + # in logs_url as that resolution key; None means "no fetchable logs". + is_actions_check = (entry.get("app") or {}).get("slug") == _GITHUB_ACTIONS_APP_SLUG + run_id = ( + self._parse_workflow_run_id(entry.get("details_url", "")) if is_actions_check else None + ) + logs_url = str(run_id) if run_id is not None else None + + raw_output = entry.get("output") or {} + output = { + k: str(raw_output.get(k) or "") + for k in ("title", "summary", "text") + if raw_output.get(k) + } + + return CheckRun( + name=entry.get("name", ""), + status=status, + conclusion=conclusion, + url=entry.get("html_url", ""), + logs_url=logs_url, + output=output, + ) + + @staticmethod + def _parse_workflow_run_id(details_url: str) -> int | None: + """Extract the workflow run id from an Actions check run's details_url. + + Actions sets details_url to + ``https://github.com/{owner}/{repo}/actions/runs/{run_id}`` (sometimes + with a trailing ``/job/{job_id}``). Returns the run id, or None when the + URL isn't in that shape -- in which case the check has no resolvable + logs. + """ + marker = "/actions/runs/" + idx = details_url.find(marker) + if idx == -1: + return None + run_segment = details_url[idx + len(marker) :].split("/", 1)[0] + return int(run_segment) if run_segment.isdigit() else None + + @_translate_provider_errors + async def get_check_logs(self, repo_ref: RepositoryRef, check: CheckRun) -> str: + """Fetch the raw logs for a check run. + + Only GitHub Actions-backed check runs have fetchable logs. For those, + get_checks stored the workflow *run* id in ``logs_url`` -- not the + check-run id, which is a different resource from the Actions job id the + logs endpoint requires. This resolves the real job by listing the run's + jobs and matching the one whose name equals this check's name, then + fetches that job's logs. + + Legacy commit-status "checks" (Prow, etc.) have no logs endpoint, so + their ``logs_url`` is None; asking for their logs raises NotFoundError + because there is genuinely nothing to fetch, rather than returning + fabricated text. + + Args: + repo_ref: Repository reference. + check: The CheckRun to fetch logs for. + + Returns: + The job log content as plain text. + + Raises: + NotFoundError: If the check has no ``logs_url`` (no Actions job backs + it), if no job in the run matches the check's name, or if GitHub + returns 404 for the job's logs (translating the provider-specific + httpx.HTTPStatusError into the neutral error so nothing above the + adapter layer touches httpx directly). + SourceControlError: If ``logs_url`` is not a numeric run id, or if + more than one job in the run shares the check's name, so the + correct logs cannot be chosen unambiguously. + """ + if not check.logs_url: + raise NotFoundError( + f"No logs available for check {check.name!r}: it is not backed by " + "a GitHub Actions job (e.g. a legacy commit-status check)." + ) + + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + + # logs_url should always be a stringified run id set by _map_check_run, + # but it survives a round-trip through the Redis queue (see + # queue/models.py) as a raw string, so a malformed or legacy-format + # entry must surface as a documented SourceControlError rather than an + # uncaught ValueError leaking out of the adapter. + try: + run_id = int(check.logs_url) + except ValueError as exc: + raise SourceControlError( + f"Check {check.name!r} has a non-numeric logs_url {check.logs_url!r}; " + "expected a GitHub Actions workflow run id." + ) from exc + jobs = await client.list_workflow_run_jobs(owner, repo, run_id) + matches = [job for job in jobs if job.get("name") == check.name] + if not matches: + raise NotFoundError( + f"No Actions job named {check.name!r} was found in workflow run " + f"{run_id} of {owner}/{repo}; cannot fetch its logs." + ) + if len(matches) > 1: + raise SourceControlError( + f"Workflow run {run_id} of {owner}/{repo} has {len(matches)} jobs named " + f"{check.name!r}; cannot unambiguously choose whose logs to fetch." + ) + + job_id = matches[0].get("id") + if job_id is None: + raise SourceControlError( + f"Actions job named {check.name!r} in workflow run {run_id} of " + f"{owner}/{repo} has no id; cannot fetch its logs." + ) + + try: + return await client.get_job_logs(owner, repo, job_id=int(job_id)) + except httpx.HTTPStatusError as exc: + if exc.response.status_code == 404: + raise NotFoundError( + f"Logs for check {check.name!r} (job {job_id}) were not " + f"found in {owner}/{repo}." + ) from exc + raise + + @_translate_provider_errors + async def get_check_artifacts( + self, repo_ref: RepositoryRef, check: CheckRun + ) -> list[tuple[str, bytes]]: + """Download the workflow run's artifacts for an Actions-backed check. + + logs_url holds the workflow run id (see _map_check_run). Non-Actions + checks have logs_url=None and no artifacts. + """ + if not check.logs_url: + return [] + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + try: + run_id = int(check.logs_url) + except ValueError as exc: + raise SourceControlError( + f"Check {check.name!r} has a non-numeric logs_url {check.logs_url!r}" + ) from exc + results: list[tuple[str, bytes]] = [] + for artifact in await client.get_run_artifacts(owner, repo, run_id): + artifact_id = artifact.get("id") + if artifact_id is None: + logger.warning( + "Skipping artifact with no id in run %s of %s/%s", run_id, owner, repo + ) + continue + name = artifact.get("name", str(artifact_id)) + zip_bytes = await client.download_artifact_zip(owner, repo, artifact_id) + results.append((name, zip_bytes)) + return results + + @_translate_provider_errors + async def get_file(self, repo_ref: RepositoryRef, path: str, ref: str) -> str: + """Fetch the content of a file at a ref. + + The client's get_file_contents returns file metadata (with base64 + ``content``) on success, or None for a 404. The protocol's return type is + a plain ``str``, so a missing file must be raised as NotFoundError rather + than papered over as an empty string. + + Args: + repo_ref: Repository reference. + path: File path in the repository. + ref: Git ref (branch, tag, or commit SHA). + + Returns: + The file content decoded to a UTF-8 string. + + Raises: + NotFoundError: If the file does not exist at the given ref. + SourceControlError: If GitHub omitted the inline base64 content (it + only inlines content for files <=1MB; larger files come back + with encoding "none" and an empty content field on an otherwise + successful response, which must not be mistaken for an empty + file). + """ + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + + result = await client.get_file_contents(owner=owner, repo=repo, path=path, ref=ref) + if result is None: + raise NotFoundError(f"File {path!r} not found at ref {ref!r} in {owner}/{repo}.") + + if result.get("encoding") != "base64": + raise SourceControlError( + f"File {path!r} at ref {ref!r} in {owner}/{repo} was not returned as inline " + f"base64 content (encoding={result.get('encoding')!r}, size={result.get('size')}) " + "-- GitHub only inlines content for files up to 1MB." + ) + + # base64.b64decode tolerates the embedded newlines GitHub's API includes + # in the encoded content field. + return base64.b64decode(result["content"]).decode() + + @_translate_provider_errors + async def put_file( + self, + repo_ref: RepositoryRef, + path: str, + content: str, + message: str, + branch: str, + ) -> None: + """Create or update a file on a branch. + + GitHub's Contents API requires the current file's ``sha`` to update an + existing file (a create-without-sha 422s if the file already exists), but + the protocol signature exposes no ``sha``. So the adapter resolves it: it + looks the file up on the target branch first, passing through the existing + ``sha`` for an update, or omitting it when the file doesn't exist yet (a + create). The client base64-encodes ``content`` internally. + + This lookup-then-write is not atomic: if the file changes on ``branch`` + between the lookup and the write (a concurrent writer, a rebase), the + ``sha`` this method passes is stale and GitHub rejects the write with a + 409, or a 422 whose error message mentions the sha mismatch. Both are + translated into ConflictError rather than left as a raw httpx error, + consistent with how ensure_write_target reports a diverged fork. A 422 + NOT related to a stale sha (branch protection, an invalid path, etc.) + propagates as the raw httpx.HTTPStatusError instead of being + mislabeled as a conflict -- see _looks_like_stale_sha_error. + + Args: + repo_ref: Repository reference. + path: File path in the repository. + content: File content (plain text). + message: Commit message. + branch: Branch to commit to. + + Raises: + ConflictError: If the file was concurrently modified between the + sha lookup and the write. + """ + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + + existing = await client.get_file_contents(owner=owner, repo=repo, path=path, ref=branch) + sha = existing["sha"] if existing is not None else None + + try: + await client.create_or_update_file( + owner=owner, + repo=repo, + path=path, + content=content, + message=message, + branch=branch, + sha=sha, + ) + except httpx.HTTPStatusError as exc: + status = exc.response.status_code + if status == 409 or (status == 422 and _looks_like_stale_sha_error(exc.response)): + raise ConflictError( + f"File {path!r} on {branch!r} in {owner}/{repo} was concurrently modified " + "since it was last read; retry with a fresh read." + ) from exc + raise + + @_translate_provider_errors + async def create_branch(self, repo_ref: RepositoryRef, name: str, base: str) -> None: + """Create a new branch from a base ref. + + The client swallows a 422 branch-already-exists into a synthesized + success, so this is idempotent from the adapter's perspective. The + protocol returns None, so the client's response dict is discarded. + + Args: + repo_ref: Repository reference. + name: New branch name. + base: Base ref to branch from. + """ + owner, repo = repo_ref.namespace.split("/", 1) + client = self._get_client() + + await client.create_branch(owner=owner, repo=repo, branch_name=name, base=base) + + @_translate_provider_errors + async def get_authenticated_identity(self, _repo_ref: RepositoryRef) -> Actor: + """Get the authenticated user's identity. + + Args: + _repo_ref: Repository reference (unused but part of the protocol) + + Returns: + Actor representing the authenticated user + """ + client = self._get_client() + user_data = await client.get_authenticated_user() + login = user_data.get("login", "") + is_bot = user_data.get("type") == "Bot" or "[bot]" in login + return Actor(login=login, is_bot=is_bot) + + async def close(self) -> None: + """Close the underlying GitHubClient's httpx.AsyncClient, if one was + ever lazily constructed.""" + if self._client is not None: + await self._client.close() diff --git a/src/forge/integrations/source_control/registry.py b/src/forge/integrations/source_control/registry.py new file mode 100644 index 000000000..492696e43 --- /dev/null +++ b/src/forge/integrations/source_control/registry.py @@ -0,0 +1,369 @@ +"""Repository/connection registry. + +Loads an optional repos.yaml-shaped config file and resolves identifiers +(explicit repository ids or provider-native namespaces) to a repository, +connection, and — once a provider has registered one — its +adapter. +""" + +import logging +import os +from collections.abc import Callable +from dataclasses import dataclass +from functools import lru_cache +from pathlib import Path +from typing import Any + +import yaml # type: ignore[import-untyped] +from pydantic import SecretStr + +from forge.config import Settings, get_settings +from forge.integrations.source_control.contracts import ( + Connection, + Provider, + RepositoryRef, + ResolvedRepository, + SourceControlProvider, +) +from forge.integrations.source_control.errors import NotFoundError, ProviderConfigError + +logger = logging.getLogger(__name__) + +IMPLICIT_GITHUB_CONNECTION_NAME = "github-default" + +AdapterFactory = Callable[[Connection], SourceControlProvider] +_ADAPTER_FACTORIES: dict[Provider, AdapterFactory] = {} + + +def register_adapter_factory(provider: Provider, factory: AdapterFactory) -> None: + """Register the adapter constructor a provider's connections resolve to. + + Called by each provider's integration package at import time (the GitHub + adapter registers itself in segment 2) — contracts.py and registry.py stay + free of provider-specific imports. + """ + _ADAPTER_FACTORIES[provider] = factory + + +@dataclass(frozen=True) +class _ImplicitConnection: + """An implicit connection paired with whether its credential is usable right now.""" + + connection: Connection + configured: bool + + +def _build_implicit_connections(settings: Settings) -> dict[Provider, _ImplicitConnection]: + """The zero-config GitHub connection every unregistered GitHub namespace resolves to. + + `configured` is checked against `settings.github_token` rather than + `os.environ`, since Settings can load GITHUB_TOKEN from .env without ever + adding it to the environment. + """ + return { + Provider.GITHUB: _ImplicitConnection( + connection=Connection( + name=IMPLICIT_GITHUB_CONNECTION_NAME, + provider=Provider.GITHUB, + base_url="https://api.github.com", + credential_env="GITHUB_TOKEN", + webhook_secret_env="GITHUB_WEBHOOK_SECRET", + ca_path=None, + allowed_namespaces=None, + ), + configured=bool(settings.github_token.get_secret_value()), + ) + } + + +class Registry: + """Resolves repository identifiers against configured connections and repositories.""" + + def __init__( + self, + connections: dict[str, Connection], + repositories: dict[str, RepositoryRef], + implicit_connections: dict[Provider, _ImplicitConnection], + ) -> None: + self._connections = connections + self._repositories = repositories + self._implicit_connections = implicit_connections + self._namespace_index: dict[tuple[Provider, str], RepositoryRef] = { + (repo.provider, repo.namespace): repo for repo in repositories.values() + } + # One adapter instance per connection for this Registry's lifetime + # (itself a process-wide singleton via get_registry()), so repeated + # resolve() calls against the same connection reuse its underlying + # HTTP client/connection pool instead of leaking a new one each time. + # Keyed by connection name, which is unique across both explicit + # repos.yaml connections and the implicit per-provider defaults. + self._adapter_cache: dict[str, SourceControlProvider] = {} + + def get_connection(self, name: str) -> Connection | None: + return self._connections.get(name) + + def get_repository(self, repository_id: str) -> RepositoryRef | None: + return self._repositories.get(repository_id) + + def resolve(self, identifier: str, provider_hint: Provider | None = None) -> ResolvedRepository: + """Resolve an explicit repository id or a provider-native namespace. + + Explicit repos.yaml ids are tried first. Anything else is treated as a + namespace under provider_hint (defaulting to GitHub, since that's the + only provider a bare namespace has ever meant). A namespace with no + explicit repositories: entry falls back to that provider's implicit + default connection; a provider with no implicit default raises + NotFoundError. + """ + repo_ref = self._repositories.get(identifier) + if repo_ref is not None: + return self._build_resolved(repo_ref, self._connections[repo_ref.connection]) + + provider = provider_hint or Provider.GITHUB + repo_ref = self._namespace_index.get((provider, identifier)) + if repo_ref is not None: + return self._build_resolved(repo_ref, self._connections[repo_ref.connection]) + + implicit_entry = self._implicit_connections.get(provider) + if implicit_entry is None: + raise NotFoundError( + f"'{identifier}' does not match a registered repository, and " + f"'{provider}' has no implicit default connection" + ) + if not implicit_entry.configured: + raise ProviderConfigError( + f"'{identifier}' resolves to the implicit '{provider}' connection, but " + f"credential_env '{implicit_entry.connection.credential_env}' is not set" + ) + + implicit_connection = implicit_entry.connection + implicit_ref = RepositoryRef( + id=identifier, + provider=provider, + connection=implicit_connection.name, + namespace=identifier, + default_branch="main", + change_request_mode="fork", + ) + return self._build_resolved(implicit_ref, implicit_connection) + + def _build_resolved( + self, repo_ref: RepositoryRef, connection: Connection + ) -> ResolvedRepository: + adapter = self._adapter_cache.get(connection.name) + if adapter is None: + factory = _ADAPTER_FACTORIES.get(connection.provider) + if factory is not None: + adapter = factory(connection) + self._adapter_cache[connection.name] = adapter + return ResolvedRepository(repo_ref=repo_ref, connection=connection, adapter=adapter) + + async def aclose(self) -> None: + """Close every adapter this Registry has cached. + + Call once at process shutdown (FastAPI lifespan, worker shutdown) -- + not per-request. Safe to call even if some/all adapters were never + actually used (their close() is a no-op in that case). + """ + for adapter in self._adapter_cache.values(): + await adapter.close() + + +def resolve_env_value(name: str, settings: Settings) -> str | None: + """Look up a named env var, preferring the matching Settings field. + + A field Settings models (e.g. GITHUB_TOKEN -> settings.github_token) must + be read through Settings rather than os.environ: BaseSettings loads .env + values directly without ever exporting them into the process environment + (see _build_implicit_connections). Names Settings doesn't model fall back + to os.environ, since repos.yaml connections can reference credentials + (e.g. for a provider without a dedicated Settings field) Settings never + claimed ownership of. + """ + field_name = name.lower() + if field_name in type(settings).model_fields: + value = getattr(settings, field_name) + if isinstance(value, SecretStr): + return value.get_secret_value() or None + return str(value) if value else None + return os.environ.get(name) + + +def _parse_connections(raw: dict[str, Any], settings: Settings) -> dict[str, Connection]: + connections: dict[str, Connection] = {} + for name, entry in raw.items(): + if name == IMPLICIT_GITHUB_CONNECTION_NAME: + raise ProviderConfigError( + f"connection '{name}' collides with the reserved implicit connection name " + f"'{IMPLICIT_GITHUB_CONNECTION_NAME}'; choose a different name" + ) + if not isinstance(entry, dict): + raise ProviderConfigError(f"connection '{name}' must be a mapping") + if "provider" not in entry: + raise ProviderConfigError(f"connection '{name}' is missing 'provider'") + try: + provider = Provider(entry["provider"]) + except ValueError as exc: + raise ProviderConfigError( + f"connection '{name}' has unknown provider '{entry['provider']}'" + ) from exc + + credential_env = entry.get("credential_env") + if not credential_env: + raise ProviderConfigError(f"connection '{name}' is missing 'credential_env'") + if not resolve_env_value(credential_env, settings): + raise ProviderConfigError( + f"connection '{name}' references credential_env '{credential_env}', " + "which is not set" + ) + + # webhook_secret_env is intentionally optional here: a connection used + # only for API operations (git push, PR creation) with no inbound + # webhook has no secret to configure. A connection that *does* receive + # webhooks but omits it fails closed at request time instead of at + # startup -- GitHubAdapter.verify_webhook rejects every delivery when + # no secret is configured, logging a warning each time. + allowed_namespaces = entry.get("allowed_namespaces") + if allowed_namespaces is not None and ( + not isinstance(allowed_namespaces, list) + or not all(isinstance(namespace, str) for namespace in allowed_namespaces) + ): + raise ProviderConfigError( + f"connection '{name}' has invalid 'allowed_namespaces': must be a list of strings" + ) + + connections[name] = Connection( + name=name, + provider=provider, + base_url=entry.get("base_url", ""), + credential_env=credential_env, + webhook_secret_env=entry.get("webhook_secret_env", ""), + ca_path=entry.get("ca_path"), + allowed_namespaces=allowed_namespaces, + ) + return connections + + +def _parse_repositories( + raw: dict[str, Any], connections: dict[str, Connection] +) -> dict[str, RepositoryRef]: + repositories: dict[str, RepositoryRef] = {} + seen_namespaces: dict[tuple[Provider, str], str] = {} + for repo_id, entry in raw.items(): + if not isinstance(entry, dict): + raise ProviderConfigError(f"repository '{repo_id}' must be a mapping") + connection_name = entry.get("connection") + if connection_name not in connections: + raise ProviderConfigError( + f"repository '{repo_id}' references unknown connection '{connection_name}'" + ) + connection = connections[connection_name] + + if "provider" not in entry: + raise ProviderConfigError(f"repository '{repo_id}' is missing 'provider'") + try: + provider = Provider(entry["provider"]) + except ValueError as exc: + raise ProviderConfigError( + f"repository '{repo_id}' has unknown provider '{entry['provider']}'" + ) from exc + + if provider != connection.provider: + raise ProviderConfigError( + f"repository '{repo_id}' has provider '{provider}' but its connection " + f"'{connection_name}' has provider '{connection.provider}'" + ) + + namespace = entry.get("namespace") + if not namespace: + raise ProviderConfigError(f"repository '{repo_id}' is missing 'namespace'") + + if ( + connection.allowed_namespaces is not None + and namespace not in connection.allowed_namespaces + ): + raise ProviderConfigError( + f"repository '{repo_id}' namespace '{namespace}' is not in connection " + f"'{connection_name}''s allowed_namespaces" + ) + + namespace_key = (provider, namespace) + if namespace_key in seen_namespaces: + raise ProviderConfigError( + f"repository '{repo_id}' and '{seen_namespaces[namespace_key]}' both use " + f"namespace '{namespace}' for provider '{provider}'" + ) + seen_namespaces[namespace_key] = repo_id + + change_request_mode = entry.get("change_request_mode", "fork") + if change_request_mode not in ("fork", "direct"): + raise ProviderConfigError( + f"repository '{repo_id}' has invalid change_request_mode '{change_request_mode}'" + ) + + repositories[repo_id] = RepositoryRef( + id=repo_id, + provider=provider, + connection=connection_name, + namespace=namespace, + default_branch=entry.get("default_branch", "main"), + change_request_mode=change_request_mode, + ) + return repositories + + +def _read_yaml_mapping(path: Path) -> Any: + try: + return yaml.safe_load(path.read_text()) or {} + except yaml.YAMLError as exc: + raise ProviderConfigError(f"{path}: invalid YAML: {exc}") from exc + + +def load_registry( + config_path: str | Path | None = None, settings: Settings | None = None +) -> Registry: + """Load and validate the repos.yaml-shaped registry config. + + Raises ProviderConfigError on any misconfiguration: unknown provider, + unknown connection reference, missing credential env var, or a + repository namespace excluded by its connection's allowed_namespaces. + A missing config file is not an error — repos.yaml is optional. + """ + settings = settings or get_settings() + path = Path(config_path) if config_path is not None else Path(settings.forge_repos_config_path) + + raw = _read_yaml_mapping(path) if path.exists() else {} + if not isinstance(raw, dict): + raise ProviderConfigError(f"{path} must contain a YAML mapping at the top level") + + connections_raw = raw.get("connections") or {} + if not isinstance(connections_raw, dict): + raise ProviderConfigError(f"{path}: 'connections' must be a mapping") + repositories_raw = raw.get("repositories") or {} + if not isinstance(repositories_raw, dict): + raise ProviderConfigError(f"{path}: 'repositories' must be a mapping") + + connections = _parse_connections(connections_raw, settings) + repositories = _parse_repositories(repositories_raw, connections) + implicit_connections = _build_implicit_connections(settings) + logger.info( + "Loaded source-control registry from %s: %d connection(s), %d repositor(y/ies)", + path, + len(connections), + len(repositories), + ) + return Registry( + connections=connections, + repositories=repositories, + implicit_connections=implicit_connections, + ) + + +@lru_cache +def get_registry() -> Registry: + """Get the cached, process-wide registry loaded from settings.forge_repos_config_path. + + Cached for the life of the process: repos.yaml edits require a restart to + take effect (see CLAUDE.md). + """ + return load_registry() diff --git a/src/forge/main.py b/src/forge/main.py index 3eb082fa7..7dd7c7b39 100644 --- a/src/forge/main.py +++ b/src/forge/main.py @@ -9,10 +9,12 @@ from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware +import forge.integrations.source_control.github # noqa: F401 (registers GitHub adapter factory) from forge import __version__ from forge.api.middleware.correlation import CorrelationIdMiddleware from forge.api.routes import github_router, health_router, jira_router, metrics_router from forge.config import get_settings +from forge.integrations.source_control.registry import get_registry from forge.observability.config import configure_tracing, shutdown_tracing from forge.orchestrator.checkpointer import close_redis_pool @@ -34,6 +36,11 @@ async def lifespan(_app: FastAPI) -> AsyncGenerator[None, None]: log_startup_banner("API Gateway") + # Load the source-control registry now so a misconfigured repos.yaml + # (unknown provider, missing credential_env, etc.) fails startup instead + # of surfacing as a 500 on the first inbound webhook. + registry = get_registry() + # Startup - initialize tracing if settings.tracing_enabled: configure_tracing( @@ -48,6 +55,7 @@ async def lifespan(_app: FastAPI) -> AsyncGenerator[None, None]: logger.info("Shutting down Forge...") if settings.tracing_enabled: await shutdown_tracing() + await registry.aclose() await close_redis_pool() diff --git a/src/forge/models/events.py b/src/forge/models/events.py index 4520f9883..bcfcfaff2 100644 --- a/src/forge/models/events.py +++ b/src/forge/models/events.py @@ -10,7 +10,7 @@ class EventSource(StrEnum): """Source of webhook events.""" JIRA = "jira" - GITHUB = "github" + SOURCE_CONTROL = "source_control" class EventStatus(StrEnum): diff --git a/src/forge/orchestrator/worker.py b/src/forge/orchestrator/worker.py index 0089b72c6..b803b4872 100644 --- a/src/forge/orchestrator/worker.py +++ b/src/forge/orchestrator/worker.py @@ -18,14 +18,23 @@ record_workflow_started, ) from forge.config import get_settings -from forge.integrations.github.client import GitHubClient -from forge.integrations.github.comment_signature import is_self_comment, resolve_bot_login +from forge.integrations.github.comment_signature import is_self_comment from forge.integrations.jira.client import JiraClient +from forge.integrations.source_control.contracts import ( + ChangeRequestState, + CheckStatus, + EventKind, + NormalizedEvent, + RepositoryRef, + Review, + ReviewState, +) +from forge.integrations.source_control.registry import get_registry from forge.models.events import EventSource from forge.models.workflow import ForgeLabel, TicketType from forge.orchestrator.checkpointer import get_checkpointer, get_ticket_from_pr_index from forge.queue.consumer import QueueConsumer -from forge.queue.models import QueueMessage +from forge.queue.models import QueueMessage, normalized_event_from_dict from forge.skills.orchestrator import ensure_skills from forge.skills.utils import extract_project_key from forge.utils.redaction import redact_secrets @@ -52,15 +61,56 @@ ) from forge.workflow.utils.review_decisions import ( decision_matches_comment, - flatten_review_threads, merge_review_decisions, ) +from forge.workflow.utils.source_control import get_adapter, identity_for logger = logging.getLogger(__name__) _CI_STAGES = ("ci_evaluator", "attempt_ci_fix", "human_review_gate") +def _flatten_review_threads(reviews: list[Review]) -> list[dict[str, Any]]: + """Return the latest comment from each non-empty review thread. + + Mirrors workflow.utils.review_decisions.flatten_review_threads, sourced + from adapter-mapped Review objects (one per thread) instead of the raw + GraphQL-shaped dicts that helper expects. + """ + return [ + { + "path": review.comments[-1].path or "", + "line": review.comments[-1].line, + "body": review.comments[-1].body, + } + for review in reviews + if review.comments + ] + + +def _reviews_to_raw_threads(reviews: list[Review]) -> list[dict[str, Any]]: + """Convert adapter-mapped Review objects (one per thread) into the raw + dict shape triage_proposal_review_threads/reply_to_proposal_decisions and + the proposal-thread diffing below expect: JSON-serializable dicts with + "thread_id"/"comments" keys, not dataclasses. + """ + return [ + { + "thread_id": review.id, + "path": review.comments[0].path if review.comments else None, + "line": review.comments[0].line if review.comments else None, + "comments": [ + { + "comment_id": int(c.id) if c.id.isdigit() else c.id, + "body": c.body, + } + for c in review.comments + ], + } + for review in reviews + ] + + def _is_workflow_errored(state: dict) -> bool: """Return True when workflow has a recorded error and is not paused for human input.""" return not state.get("is_paused") and state.get("last_error") is not None @@ -149,21 +199,34 @@ def __init__( self._shutdown_event = asyncio.Event() self._checkpointer = None self._compiled_workflows: dict[str, Any] = {} # Cache compiled workflows by name - self._forge_github_login: str | None = None + # Keyed by connection name -- a different connection may authenticate + # as a different bot identity, so a single process-wide login is wrong + # once more than the default connection is configured. + self._forge_github_logins: dict[str, str] = {} + + def _deserialize_event(self, message: QueueMessage) -> NormalizedEvent | None: + """Reconstruct the typed NormalizedEvent a source-control message carries. + + Returns None for Jira messages (which never set normalized_event) or for + a source-control message that predates this field for some reason (e.g. + a backlog entry queued before this field existed). Callers currently + treat None as "no match" / "nothing to detect" rather than falling back + to raw-payload handling -- there is no fallback path implemented. + """ + if message.normalized_event is None: + return None + return normalized_event_from_dict(message.normalized_event) - async def _get_forge_github_login(self) -> str: - """Resolve and cache the authenticated Forge GitHub login.""" - cached = getattr(self, "_forge_github_login", None) + async def _get_forge_github_login(self, repo_ref: RepositoryRef) -> str: + """Resolve and cache the authenticated Forge identity for this connection.""" + cached = self._forge_github_logins.get(repo_ref.connection) if cached: return cached - github = GitHubClient() - try: - login = await resolve_bot_login(github) - finally: - await github.close() - if login: - self._forge_github_login = login - return login + _, adapter = get_adapter(repo_ref.namespace) + identity = await adapter.get_authenticated_identity(repo_ref) + if identity.login: + self._forge_github_logins[repo_ref.connection] = identity.login + return identity.login async def _handle_terminal_failure(self, message: QueueMessage, error: str) -> None: """Post one Jira comment after queue retries are exhausted.""" @@ -204,8 +267,8 @@ async def _handle_jira_event(self, message: QueueMessage) -> None: """ await self._process_workflow(message) - async def _handle_github_event(self, message: QueueMessage) -> None: - """Handle a GitHub webhook event. + async def _handle_source_control_event(self, message: QueueMessage) -> None: + """Handle a source-control webhook event. Args: message: The queue message to process. @@ -214,7 +277,7 @@ async def _handle_github_event(self, message: QueueMessage) -> None: message = await self._resolve_ticket_from_pr_index(message) if not message.ticket_key: logger.info( - f"Dropping GitHub event {message.event_id}: " + f"Dropping source-control event {message.event_id}: " "no ticket key in message and PR URL not found in Redis index" ) return @@ -281,38 +344,40 @@ async def _resolve_ticket_from_pr_index(self, message: QueueMessage) -> QueueMes return message def _is_prd_pr_event(self, message: QueueMessage, current_state: dict[str, Any]) -> bool: - """Check if a GitHub event targets the PRD proposals PR.""" - if message.source != EventSource.GITHUB: + """Check if a source-control event targets the PRD proposals PR.""" + if message.source != EventSource.SOURCE_CONTROL: return False prd_pr_number = current_state.get("prd_pr_number") prd_pr_repo = current_state.get("prd_pr_repo") if not prd_pr_number or not prd_pr_repo: return False - payload = message.payload - repo_full = payload.get("repository", {}).get("full_name", "") - event_pr_number = payload.get("pull_request", {}).get("number") or payload.get( - "issue", {} - ).get("number") + event = self._deserialize_event(message) + if event is None or event.change_request is None: + return False - return repo_full == prd_pr_repo and event_pr_number == prd_pr_number + return ( + event.repo_ref.namespace == prd_pr_repo + and event.change_request.identity.native_id == prd_pr_number + ) def _is_spec_pr_event(self, message: QueueMessage, current_state: dict[str, Any]) -> bool: - """Check if a GitHub event targets the spec proposals PR.""" - if message.source != EventSource.GITHUB: + """Check if a source-control event targets the spec proposals PR.""" + if message.source != EventSource.SOURCE_CONTROL: return False spec_pr_number = current_state.get("spec_pr_number") spec_pr_repo = current_state.get("spec_pr_repo") if not spec_pr_number or not spec_pr_repo: return False - payload = message.payload - repo_full = payload.get("repository", {}).get("full_name", "") - event_pr_number = payload.get("pull_request", {}).get("number") or payload.get( - "issue", {} - ).get("number") + event = self._deserialize_event(message) + if event is None or event.change_request is None: + return False - return repo_full == spec_pr_repo and event_pr_number == spec_pr_number + return ( + event.repo_ref.namespace == spec_pr_repo + and event.change_request.identity.native_id == spec_pr_number + ) async def _process_workflow(self, message: QueueMessage) -> None: """Process a message through the workflow. @@ -555,8 +620,9 @@ 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) + event_obj = self._deserialize_event(message) + current_state = activate_pull_request_for_event(current_state, event_obj) + targets_implementation_pr = event_targets_pull_request(current_state, event_obj) changelog = payload.get("changelog", {}) comment = payload.get("comment", {}) @@ -584,27 +650,34 @@ async def _handle_resume_event( # Preserve unrelated contested threads and re-run review analysis so any # newly accepted item can proceed without globally clearing objections. if ( - message.source == EventSource.GITHUB - and "pull_request_review_comment" in message.event_type + event_obj is not None + and event_obj.kind == EventKind.COMMENT_CREATED + and event_obj.comment is not None + and event_obj.comment.path is not None and current_node == "review_response_gate" and current_state.get("is_paused", True) ): - reply = payload.get("comment", {}) - replied_to = reply.get("in_reply_to_id") - sender_login = payload.get("sender", {}).get("login", "") + reply = event_obj.comment + sender_login = event_obj.actor.login if sender_login: - forge_login = await self._get_forge_github_login() + forge_login = await self._get_forge_github_login(event_obj.repo_ref) settings = get_settings() forge_bot_comment_prefix = settings.forge_bot_comment_prefix if is_self_comment( sender_login=sender_login, - comment_body=reply.get("body", ""), + comment_body=reply.body, bot_login=forge_login, prefix=forge_bot_comment_prefix, ): logger.debug("Ignoring Forge's own inline review comment") return current_state - if replied_to: + in_reply_to_raw = reply.in_reply_to + replied_to = ( + int(in_reply_to_raw) + if in_reply_to_raw is not None and in_reply_to_raw.isdigit() + else None + ) + if replied_to is not None: contested = current_state.get("contested_comments", []) remaining = [ item for item in contested if not decision_matches_comment(item, replied_to) @@ -613,7 +686,7 @@ async def _handle_resume_event( **current_state, "is_paused": False, "revision_requested": True, - "feedback_comment": reply.get("body", ""), + "feedback_comment": reply.body, "contested_comments": remaining, "context": { **current_state.get("context", {}), @@ -622,59 +695,60 @@ async def _handle_resume_event( "review_thread_comment_id": replied_to, }, } + own_id = int(reply.id) if reply.id and reply.id.isdigit() else None return { **current_state, "is_paused": False, "revision_requested": True, - "feedback_comment": reply.get("body", ""), + "feedback_comment": reply.body, "context": { **current_state.get("context", {}), "resume_event": message.event_type, "payload": payload, - "review_thread_comment_id": reply.get("id"), + "review_thread_comment_id": own_id, }, } - # GitHub check_run/check_suite events are the explicit signal for wait_for_ci_gate. - # They don't carry Jira labels or comments, so handle them before the label loop. - # For check_suite and check_run events (both real GitHub webhooks and poller - # forwarded ones), only wake up CI evaluation when the suite is completed. - # 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. - event = message.event_type - is_check_event = "check_suite" in event or "check_run" in event - if message.source == EventSource.GITHUB and ( + is_check_event = event_obj is not None and event_obj.kind == EventKind.CHECK_UPDATED + if event_obj is not None and ( current_node == "ci_evaluator" or (targets_implementation_pr and is_check_event) ): if is_check_event: - suite_status = payload.get("check_suite", {}).get("status") or payload.get( - "check_run", {} - ).get("check_suite", {}).get("status") - if suite_status and suite_status != "completed": + suite_status = event_obj.check_suite_status + if suite_status and suite_status != CheckStatus.COMPLETED: logger.info( - f"Ignoring {event} for {message.ticket_key}: " + f"Ignoring {message.event_type} for {message.ticket_key}: " f"check_suite not yet completed (status={suite_status!r})" ) else: is_ci_webhook = True - logger.info(f"Detected GitHub CI webhook signal for {current_node}") - elif ( - "issue_comment" not in event - and "pull_request_review" not in event - and payload.get("pull_request", {}).get("merged") is not True + logger.info(f"Detected source-control CI webhook signal for {current_node}") + elif not ( + event_obj.kind + in (EventKind.COMMENT_CREATED, EventKind.REVIEW_SUBMITTED, EventKind.UNKNOWN) + or ( + event_obj.change_request + and event_obj.change_request.state == ChangeRequestState.MERGED + ) ): is_ci_webhook = True - logger.info(f"Detected GitHub CI webhook signal for {current_node}") + logger.info(f"Detected source-control CI webhook signal for {current_node}") # GitHub issue_comment events: detect /forge skip-gate and /forge unskip-gate # commands posted as PR comments. - if message.source == EventSource.GITHUB and "issue_comment" in message.event_type: - gh_comment_body = payload.get("comment", {}).get("body", "").strip() - repo_full = payload.get("repository", {}).get("full_name", "") - pr_number = payload.get("issue", {}).get("number") - sender = payload.get("sender", {}).get("login", "") - _owner, _, _repo = repo_full.partition("/") + if ( + event_obj is not None + and event_obj.kind == EventKind.COMMENT_CREATED + and event_obj.comment is not None + and event_obj.comment.path is None + ): + gh_comment_body = (event_obj.comment.body or "").strip() + repo_full = event_obj.repo_ref.namespace + native_id = ( + event_obj.change_request.identity.native_id if event_obj.change_request else None + ) + pr_number = int(native_id) if native_id is not None else None + sender = event_obj.actor.login skip_prefix = "/forge skip-gate" unskip_prefix = "/forge unskip-gate" @@ -688,8 +762,7 @@ async def _handle_resume_event( logger.info(f"CI gate skip added for {message.ticket_key}: '{check_name}'") await self._post_skip_gate_feedback( ticket_key=message.ticket_key, - owner=_owner, - repo=_repo, + repo_ref=event_obj.repo_ref, pr_number=pr_number, check_name=check_name, sender=sender, @@ -712,8 +785,7 @@ async def _handle_resume_event( logger.info(f"CI gate skip removed for {message.ticket_key}: '{check_name}'") await self._post_skip_gate_feedback( ticket_key=message.ticket_key, - owner=_owner, - repo=_repo, + repo_ref=event_obj.repo_ref, pr_number=pr_number, check_name=check_name, sender=sender, @@ -738,8 +810,7 @@ async def _handle_resume_event( logger.info(f"Detected /forge rebase for {message.ticket_key}") await self._post_rebase_feedback( ticket_key=message.ticket_key, - owner=_owner, - repo=_repo, + repo_ref=event_obj.repo_ref, pr_number=pr_number, sender=sender, ) @@ -996,29 +1067,35 @@ async def _handle_resume_event( # A human reply to a proposal review thread resumes only that thread's # feedback. Forge-authored replies are informational and must not loop. if ( - message.source == EventSource.GITHUB - and "pull_request_review_comment" in message.event_type + event_obj is not None + and event_obj.kind == EventKind.COMMENT_CREATED + and event_obj.comment is not None + and event_obj.comment.path is not None ): + reply = event_obj.comment + in_reply_to_raw = reply.in_reply_to + replied_to = ( + int(in_reply_to_raw) + if in_reply_to_raw is not None and in_reply_to_raw.isdigit() + else None + ) is_proposal_reply = ( self._is_prd_pr_event(message, current_state) and current_node in _PRD_GATE_NODES ) or ( self._is_spec_pr_event(message, current_state) and current_node in _SPEC_GATE_NODES ) - reply = payload.get("comment", {}) - replied_to = reply.get("in_reply_to_id") - if is_proposal_reply: - sender_login = payload.get("sender", {}).get("login", "") - if sender_login: - forge_login = await self._get_forge_github_login() - settings = get_settings() - forge_bot_comment_prefix = settings.forge_bot_comment_prefix - if is_self_comment( - sender_login=sender_login, - comment_body=reply.get("body", ""), - bot_login=forge_login, - prefix=forge_bot_comment_prefix, - ): - return current_state + sender_login = event_obj.actor.login + if is_proposal_reply and sender_login: + forge_login = await self._get_forge_github_login(event_obj.repo_ref) + settings = get_settings() + forge_bot_comment_prefix = settings.forge_bot_comment_prefix + if is_self_comment( + sender_login=sender_login, + comment_body=reply.body, + bot_login=forge_login, + prefix=forge_bot_comment_prefix, + ): + return current_state if is_proposal_reply and replied_to: previous = current_state.get("proposal_review_decisions", []) matching = next( @@ -1026,11 +1103,16 @@ async def _handle_resume_event( None, ) if matching: - reply_body = reply.get("body", "").strip() + reply_body = reply.body.strip() + reply_comment_id = int(reply.id) if reply.id.isdigit() else None decisions = [ { **item, - "comment_id": reply.get("id", item.get("comment_id")), + "comment_id": ( + reply_comment_id + if reply_comment_id is not None + else item.get("comment_id") + ), "disposition": "accept", "feedback": reply_body, "status": "pending", @@ -1053,59 +1135,69 @@ async def _handle_resume_event( replied_to, ) elif is_proposal_reply: - body = reply.get("body", "").strip() - comment_id = reply.get("id") - if body and isinstance(comment_id, int): + body = reply.body.strip() + if body and reply.id.isdigit(): + comment_id = int(reply.id) proposal_review_threads = [ { "thread_id": f"comment-{comment_id}", - "path": reply.get("path", ""), - "line": reply.get("line") or reply.get("original_line"), + "path": reply.path or "", + "line": reply.line, "comments": [ { "comment_id": comment_id, "body": body, "author": sender_login, - "commit_sha": reply.get("commit_id", ""), + "commit_sha": event_obj.raw.get("comment", {}).get( + "commit_id", "" + ), } ], } ] is_rejected = True feedback = body + else: + logger.warning( + "Dropping proposal reply with empty body or non-numeric " + f"comment id (id={reply.id!r}) for {message.ticket_key}" + ) # GitHub events targeting the PRD proposals PR — handled at prd_approval_gate. # Merge = approval. Review with feedback = revision. Comment = feedback/question. if self._is_prd_pr_event(message, current_state) and current_node in _PRD_GATE_NODES: - event = message.event_type - - if "pull_request_review" in event: - review = payload.get("review", {}) - review_state = review.get("state", "").lower() - review_body = review.get("body", "") or "" + if ( + event_obj is not None + and event_obj.kind == EventKind.REVIEW_SUBMITTED + and event_obj.review is not None + ): + pr_review = event_obj.review # Merge-only approval: review approval is intentionally ignored - if review_state in ("changes_requested", "commented"): - repo_full = payload.get("repository", {}).get("full_name", "") - pr_number = payload.get("pull_request", {}).get("number") + if pr_review.state in (ReviewState.CHANGES_REQUESTED, ReviewState.COMMENTED): + repo_full = event_obj.repo_ref.namespace + native_id = ( + event_obj.change_request.identity.native_id + if event_obj.change_request + else None + ) + pr_number = int(native_id) if native_id is not None else None inline_comments: list[dict[str, Any]] = [] if repo_full and pr_number: - _owner, _repo = repo_full.split("/", 1) - gh = GitHubClient() - try: - proposal_review_threads = await gh.get_pull_request_review_threads( - _owner, _repo, pr_number - ) - inline_comments = flatten_review_threads(proposal_review_threads) - finally: - await gh.close() + _repo_ref_obj, _adapter = get_adapter(repo_full) + _identity = identity_for(_repo_ref_obj, pr_number) + _reviews = await _adapter.get_review_thread_comments( + _repo_ref_obj, _identity + ) + proposal_review_threads = _reviews_to_raw_threads(_reviews) + inline_comments = _flatten_review_threads(_reviews) parts = [] - if review_body.strip(): - parts.append(review_body.strip()) + if pr_review.body.strip(): + parts.append(pr_review.body.strip()) if inline_comments: inline_text = "\n\n".join( - f"**{c['path']}** (line {c.get('line') or c.get('original_line', '?')}):\n{c['body']}" + f"**{c['path']}** (line {c.get('line') or '?'}):\n{c['body']}" for c in inline_comments ) parts.append(f"Inline comments:\n{inline_text}") @@ -1114,18 +1206,22 @@ async def _handle_resume_event( feedback = "\n\n".join(parts) is_rejected = True logger.info( - f"PRD PR review ({review_state}) for {message.ticket_key}: " - f"body={'yes' if review_body.strip() else 'no'}, " + f"PRD PR review ({pr_review.state.value}) for {message.ticket_key}: " + f"body={'yes' if pr_review.body.strip() else 'no'}, " f"inline={len(inline_comments)}" ) else: logger.info( - f"PRD PR review ({review_state}) for {message.ticket_key} " + f"PRD PR review ({pr_review.state.value}) for {message.ticket_key} " "with no content — ignoring" ) return current_state - elif "pull_request" in event and payload.get("pull_request", {}).get("merged") is True: + elif ( + event_obj is not None + and event_obj.change_request is not None + and event_obj.change_request.state == ChangeRequestState.MERGED + ): is_approved = True pr_merged = True logger.info(f"PRD PR merged for {message.ticket_key}") @@ -1141,19 +1237,18 @@ async def _handle_resume_event( finally: await jira.close() - elif "issue_comment" in event: - gh_comment = payload.get("comment", {}) - comment_body = gh_comment.get("body", "").strip() - sender_login = payload.get("sender", {}).get("login", "") + elif ( + event_obj is not None + and event_obj.kind == EventKind.COMMENT_CREATED + and event_obj.comment is not None + and event_obj.comment.path is None + ): + comment_body = (event_obj.comment.body or "").strip() + sender_login = event_obj.actor.login if comment_body and sender_login: # Skip self-comments - gh = GitHubClient() - try: - forge_user = await gh.get_authenticated_user() - forge_login = forge_user.get("login", "") - finally: - await gh.close() + forge_login = await self._get_forge_github_login(event_obj.repo_ref) settings = get_settings() forge_bot_comment_prefix = settings.forge_bot_comment_prefix @@ -1187,34 +1282,37 @@ async def _handle_resume_event( # GitHub events targeting the spec proposals PR — same pattern as PRD PR. if self._is_spec_pr_event(message, current_state) and current_node in _SPEC_GATE_NODES: - event = message.event_type - - if "pull_request_review" in event: - review = payload.get("review", {}) - review_state = review.get("state", "").lower() - review_body = review.get("body", "") or "" - - if review_state in ("changes_requested", "commented"): - repo_full = payload.get("repository", {}).get("full_name", "") - pr_number = payload.get("pull_request", {}).get("number") + if ( + event_obj is not None + and event_obj.kind == EventKind.REVIEW_SUBMITTED + and event_obj.review is not None + ): + pr_review = event_obj.review + + if pr_review.state in (ReviewState.CHANGES_REQUESTED, ReviewState.COMMENTED): + repo_full = event_obj.repo_ref.namespace + native_id = ( + event_obj.change_request.identity.native_id + if event_obj.change_request + else None + ) + pr_number = int(native_id) if native_id is not None else None inline_comments: list[dict[str, Any]] = [] if repo_full and pr_number: - _owner, _repo = repo_full.split("/", 1) - gh = GitHubClient() - try: - proposal_review_threads = await gh.get_pull_request_review_threads( - _owner, _repo, pr_number - ) - inline_comments = flatten_review_threads(proposal_review_threads) - finally: - await gh.close() + _repo_ref_obj, _adapter = get_adapter(repo_full) + _identity = identity_for(_repo_ref_obj, pr_number) + _reviews = await _adapter.get_review_thread_comments( + _repo_ref_obj, _identity + ) + proposal_review_threads = _reviews_to_raw_threads(_reviews) + inline_comments = _flatten_review_threads(_reviews) parts = [] - if review_body.strip(): - parts.append(review_body.strip()) + if pr_review.body.strip(): + parts.append(pr_review.body.strip()) if inline_comments: inline_text = "\n\n".join( - f"**{c['path']}** (line {c.get('line') or c.get('original_line', '?')}):\n{c['body']}" + f"**{c['path']}** (line {c.get('line') or '?'}):\n{c['body']}" for c in inline_comments ) parts.append(f"Inline comments:\n{inline_text}") @@ -1223,18 +1321,22 @@ async def _handle_resume_event( feedback = "\n\n".join(parts) is_rejected = True logger.info( - f"Spec PR review ({review_state}) for {message.ticket_key}: " - f"body={'yes' if review_body.strip() else 'no'}, " + f"Spec PR review ({pr_review.state.value}) for {message.ticket_key}: " + f"body={'yes' if pr_review.body.strip() else 'no'}, " f"inline={len(inline_comments)}" ) else: logger.info( - f"Spec PR review ({review_state}) for {message.ticket_key} " + f"Spec PR review ({pr_review.state.value}) for {message.ticket_key} " "with no content — ignoring" ) return current_state - elif "pull_request" in event and payload.get("pull_request", {}).get("merged") is True: + elif ( + event_obj is not None + and event_obj.change_request is not None + and event_obj.change_request.state == ChangeRequestState.MERGED + ): is_approved = True pr_merged = True logger.info(f"Spec PR merged for {message.ticket_key}") @@ -1280,18 +1382,17 @@ async def _handle_resume_event( finally: await jira.close() - elif "issue_comment" in event: - gh_comment = payload.get("comment", {}) - comment_body = gh_comment.get("body", "").strip() - sender_login = payload.get("sender", {}).get("login", "") + elif ( + event_obj is not None + and event_obj.kind == EventKind.COMMENT_CREATED + and event_obj.comment is not None + and event_obj.comment.path is None + ): + comment_body = (event_obj.comment.body or "").strip() + sender_login = event_obj.actor.login if comment_body and sender_login: - gh = GitHubClient() - try: - forge_user = await gh.get_authenticated_user() - forge_login = forge_user.get("login", "") - finally: - await gh.close() + forge_login = await self._get_forge_github_login(event_obj.repo_ref) settings = get_settings() forge_bot_comment_prefix = settings.forge_bot_comment_prefix @@ -1442,62 +1543,65 @@ async def _handle_resume_event( # GitHub pull_request_review events — handled when paused at human_review_gate or review_response_gate. # A review submission is the primary signal for the human review stage. if ( - message.source == EventSource.GITHUB - and "pull_request_review" in message.event_type + event_obj is not None + and event_obj.kind == EventKind.REVIEW_SUBMITTED + and event_obj.review is not None and (current_node in _REVIEW_GATES or targets_implementation_pr) - and current_state.get("is_paused", True) + and (current_state.get("is_paused", True) or current_state.get("pending_ci_event")) ): - sender_login = payload.get("sender", {}).get("login", "") - review = payload.get("review", {}) or {} - review_body = review.get("body", "") or "" + review = event_obj.review + sender_login = review.author if sender_login: - forge_login = await self._get_forge_github_login() + forge_login = await self._get_forge_github_login(event_obj.repo_ref) settings = get_settings() forge_bot_comment_prefix = settings.forge_bot_comment_prefix if is_self_comment( sender_login=sender_login, - comment_body=review_body, + comment_body=review.body, bot_login=forge_login, prefix=forge_bot_comment_prefix, ): logger.debug("Ignoring Forge's own pull request review") return current_state - review_state = review.get("state", "").lower() - - if review_state == "approved": + if review.state == ReviewState.APPROVED: 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"): + elif review.state in (ReviewState.CHANGES_REQUESTED, ReviewState.COMMENTED): # Always fetch inline comments so the agent gets the full picture, # regardless of whether a summary body is also present. - repo_full = payload.get("repository", {}).get("full_name", "") - pr_number = payload.get("pull_request", {}).get("number") + repo_full = event_obj.repo_ref.namespace + pr_number = ( + event_obj.change_request.identity.native_id + if event_obj.change_request + else None + ) inline_comments = [] if repo_full and pr_number: - owner, repo_name = repo_full.split("/", 1) - gh = GitHubClient() - try: - review_id = review.get("id") - if review_id: - inline_comments = await gh.get_review_comments( - owner, repo_name, pr_number, review_id - ) - else: - inline_comments = await gh.get_pull_request_review_comments( - owner, repo_name, pr_number - ) - finally: - await gh.close() + _repo_ref_obj, _adapter = get_adapter(repo_full) + _identity = identity_for(_repo_ref_obj, pr_number) + review_id = int(review.id) if review.id else None + if review_id: + review_comments = await _adapter.get_review_comments_for_submission( + _repo_ref_obj, _identity, str(review_id) + ) + else: + threads = await _adapter.get_review_thread_comments( + _repo_ref_obj, _identity + ) + review_comments = [c for thread in threads for c in thread.comments] + inline_comments = [ + {"path": c.path, "line": c.line, "body": c.body} for c in review_comments + ] parts = [] - if review_body.strip(): - parts.append(review_body.strip()) + if review.body.strip(): + parts.append(review.body.strip()) if inline_comments: inline_text = "\n\n".join( - f"**{c['path']}** (line {c.get('line') or c.get('original_line') or c.get('position') or '?'}):\n{c['body']}" + f"**{c['path']}** (line {c.get('line') or '?'}):\n{c['body']}" for c in inline_comments ) parts.append(f"Inline comments:\n{inline_text}") @@ -1506,22 +1610,22 @@ async def _handle_resume_event( feedback = "\n\n".join(parts) is_rejected = True logger.info( - f"Detected PR review ({review_state}) for {message.ticket_key}: " - f"body={'yes' if review_body.strip() else 'no'}, " + f"Detected PR review ({review.state.value}) for {message.ticket_key}: " + f"body={'yes' if review.body.strip() else 'no'}, " f"inline comments={len(inline_comments)}" ) else: logger.info( - f"Detected PR review ({review_state}) for {message.ticket_key} " + f"Detected PR review ({review.state.value}) for {message.ticket_key} " f"with no body and no inline comments — ignoring" ) return current_state # GitHub pull_request:closed + merged — PR was actually merged if ( - message.source == EventSource.GITHUB - and "pull_request" in message.event_type - and payload.get("pull_request", {}).get("merged") is True + event_obj is not None + and event_obj.change_request is not None + and event_obj.change_request.state == ChangeRequestState.MERGED and (current_node in _REVIEW_GATES or targets_implementation_pr) ): is_approved = True @@ -1543,7 +1647,7 @@ async def _handle_resume_event( if targets_implementation_pr and is_ci_webhook and current_node != "human_review_gate": updated_state["current_node"] = "ci_evaluator" elif targets_implementation_pr and ( - "pull_request_review" in message.event_type or pr_merged + (event_obj is not None and event_obj.kind == EventKind.REVIEW_SUBMITTED) or pr_merged ): updated_state["current_node"] = "human_review_gate" @@ -1670,7 +1774,7 @@ async def _handle_resume_event( updated_state["human_review_status"] = "approved" if pr_merged: updated_state["pr_merged"] = True - if event_targets_pull_request(updated_state, payload): + if event_targets_pull_request(updated_state, event_obj): 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"]: @@ -1774,6 +1878,7 @@ async def _handle_resume_event( "ci_evaluator", "attempt_ci_fix", "human_review_gate", + "review_response_gate", ) if ( not current_state.get("is_paused", True) @@ -1886,8 +1991,7 @@ def _walk(nodes: list[dict]) -> None: async def _post_skip_gate_feedback( self, ticket_key: str, - owner: str, - repo: str, + repo_ref: RepositoryRef, pr_number: int | None, check_name: str, sender: str, @@ -1897,15 +2001,14 @@ async def _post_skip_gate_feedback( Args: ticket_key: Jira ticket key for the audit comment. - owner: Repository owner. - repo: Repository name. + repo_ref: Repository reference the PR belongs to. pr_number: Pull request number. check_name: The check name that was skipped or unskipped. sender: GitHub login of the user who issued the command. action: "skip" or "unskip". """ try: - github = GitHubClient() + _, adapter = get_adapter(repo_ref.namespace) jira = JiraClient() try: if action == "skip": @@ -1934,10 +2037,10 @@ async def _post_skip_gate_feedback( ) if pr_number: - await github.create_issue_comment(owner, repo, pr_number, gh_comment) + identity = identity_for(repo_ref, pr_number) + await adapter.create_comment(repo_ref, identity, gh_comment) await post_status_comment(jira, ticket_key, jira_comment) finally: - await github.close() await jira.close() except Exception as e: logger.warning(f"Failed to post skip-gate feedback: {e}") @@ -1945,14 +2048,13 @@ async def _post_skip_gate_feedback( async def _post_rebase_feedback( self, ticket_key: str, - owner: str, - repo: str, + repo_ref: RepositoryRef, pr_number: int | None, sender: str, ) -> None: """Post feedback for a /forge rebase command.""" try: - github = GitHubClient() + _, adapter = get_adapter(repo_ref.namespace) jira = JiraClient() try: gh_comment = ( @@ -1964,10 +2066,10 @@ async def _post_rebase_feedback( f"Rebase triggered via `/forge rebase` on PR #{pr_number} by {sender}." ) if pr_number: - await github.create_issue_comment(owner, repo, pr_number, gh_comment) + identity = identity_for(repo_ref, pr_number) + await adapter.create_comment(repo_ref, identity, gh_comment) await post_status_comment(jira, ticket_key, jira_comment) finally: - await github.close() await jira.close() except Exception as e: logger.warning(f"Failed to post rebase feedback: {e}") @@ -2198,7 +2300,9 @@ async def start(self) -> None: # Register handlers self.consumer.register_handler(EventSource.JIRA, self._handle_jira_event) - self.consumer.register_handler(EventSource.GITHUB, self._handle_github_event) + self.consumer.register_handler( + EventSource.SOURCE_CONTROL, self._handle_source_control_event + ) try: await self.consumer.start() @@ -2206,6 +2310,7 @@ async def start(self) -> None: pass finally: await self.consumer.stop() + await get_registry().aclose() logger.info("Worker shut down gracefully") def _handle_shutdown(self) -> None: diff --git a/src/forge/queue/consumer.py b/src/forge/queue/consumer.py index 1a95fd9fd..8a2098df3 100644 --- a/src/forge/queue/consumer.py +++ b/src/forge/queue/consumer.py @@ -15,7 +15,7 @@ from forge.models.events import EventSource from forge.orchestrator.checkpointer import get_redis_client from forge.queue.models import QueueMessage -from forge.queue.producer import GITHUB_STREAM, JIRA_STREAM +from forge.queue.producer import JIRA_STREAM, SOURCE_CONTROL_STREAM from forge.queue.retry import RETRY_CLAIM_RENEW_SECONDS, RetryEntry, RetryQueue logger = logging.getLogger(__name__) @@ -150,7 +150,7 @@ async def _ensure_consumer_groups(self) -> None: """Ensure consumer groups exist for all streams.""" redis_client = await self._get_redis() - for stream in [JIRA_STREAM, GITHUB_STREAM]: + for stream in [JIRA_STREAM, SOURCE_CONTROL_STREAM]: try: await redis_client.xgroup_create(stream, CONSUMER_GROUP, id="0", mkstream=True) logger.info(f"Created consumer group {CONSUMER_GROUP} for {stream}") @@ -407,7 +407,7 @@ async def _process_due_retries_once(self) -> None: entries = await self._retry_queue.claim_due_messages() for entry in entries: retry_stream = ( - JIRA_STREAM if entry.message.source == EventSource.JIRA else GITHUB_STREAM + JIRA_STREAM if entry.message.source == EventSource.JIRA else SOURCE_CONTROL_STREAM ) try: async with self._renew_retry_claim(entry): @@ -488,8 +488,15 @@ async def start(self) -> None: tasks = [] if EventSource.JIRA in self._handlers: tasks.append(self._consume_stream(JIRA_STREAM, EventSource.JIRA)) - if EventSource.GITHUB in self._handlers: - tasks.append(self._consume_stream(GITHUB_STREAM, EventSource.GITHUB)) + if EventSource.SOURCE_CONTROL in self._handlers: + tasks.append(self._consume_stream(SOURCE_CONTROL_STREAM, EventSource.SOURCE_CONTROL)) + # LEGACY_SOURCE_CONTROL_STREAM (the pre-rename "forge:events:github") + # is intentionally not auto-consumed: those entries predate the + # NormalizedEvent/adapter cutover and have no normalized_event to + # deserialize, so the handler could only silently no-op and ack them + # -- discarding whatever CI/review/merge signal they carried instead + # of processing it. health_check reports its depth so a nonzero + # backlog is visible for a deliberate, out-of-band migration. if tasks: tasks.append(self._process_retry_queue()) diff --git a/src/forge/queue/models.py b/src/forge/queue/models.py index 3a9d7d604..d29372d73 100644 --- a/src/forge/queue/models.py +++ b/src/forge/queue/models.py @@ -6,8 +6,35 @@ from datetime import datetime from typing import Any +from forge.integrations.source_control.contracts import ( + Actor, + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + CheckConclusion, + CheckRun, + CheckStatus, + EventKind, + NormalizedEvent, + Provider, + RepositoryRef, + Review, + ReviewComment, + ReviewState, +) from forge.models.events import EventSource +# EventSource.SOURCE_CONTROL's value was renamed from "github" to +# "source_control". Retry/DLQ entries and unconsumed stream messages +# persisted before the rename still carry the old value in Redis; map it +# forward so they keep deserializing instead of raising ValueError. +_LEGACY_SOURCE_VALUES: dict[str, EventSource] = {"github": EventSource.SOURCE_CONTROL} + + +def _parse_event_source(value: str) -> EventSource: + legacy = _LEGACY_SOURCE_VALUES.get(value) + return legacy if legacy is not None else EventSource(value) + @dataclass class QueueMessage: @@ -19,6 +46,7 @@ class QueueMessage: event_type: str ticket_key: str payload: dict[str, Any] = field(default_factory=dict) + normalized_event: dict[str, Any] | None = None timestamp: datetime = field(default_factory=datetime.utcnow) retry_count: int = 0 @@ -34,6 +62,9 @@ def to_dict(self) -> dict[str, str]: "event_type": self.event_type, "ticket_key": self.ticket_key, "payload": json.dumps(self.payload), + "normalized_event": ( + json.dumps(self.normalized_event) if self.normalized_event is not None else "" + ), "timestamp": self.timestamp.isoformat(), "retry_count": str(self.retry_count), } @@ -49,13 +80,15 @@ def from_redis(cls, message_id: str, data: dict[str, str]) -> "QueueMessage": Returns: Populated QueueMessage instance. """ + normalized_event_raw = data.get("normalized_event", "") return cls( message_id=message_id, event_id=data.get("event_id", ""), - source=EventSource(data.get("source", "jira")), + source=_parse_event_source(data.get("source", "jira")), event_type=data.get("event_type", ""), ticket_key=data.get("ticket_key", ""), payload=json.loads(data.get("payload", "{}")), + normalized_event=json.loads(normalized_event_raw) if normalized_event_raw else None, timestamp=datetime.fromisoformat(data.get("timestamp", datetime.utcnow().isoformat())), retry_count=int(data.get("retry_count", "0")), ) @@ -63,3 +96,172 @@ def from_redis(cls, message_id: str, data: dict[str, str]) -> "QueueMessage": def increment_retry(self) -> "QueueMessage": """Return a new message with incremented retry count.""" return dataclass_replace(self, retry_count=self.retry_count + 1) + + +def normalized_event_to_dict(event: NormalizedEvent) -> dict[str, Any]: + """Serialize a NormalizedEvent for queue transport.""" + return { + "id": event.id, + "kind": event.kind.value, + "repo_ref": { + "id": event.repo_ref.id, + "provider": event.repo_ref.provider.value, + "connection": event.repo_ref.connection, + "namespace": event.repo_ref.namespace, + "default_branch": event.repo_ref.default_branch, + "change_request_mode": event.repo_ref.change_request_mode, + }, + "actor": {"login": event.actor.login, "is_bot": event.actor.is_bot}, + "received_at": event.received_at.isoformat(), + "change_request": ( + { + "identity": { + "connection": event.change_request.identity.connection, + "repository_id": event.change_request.identity.repository_id, + "native_id": event.change_request.identity.native_id, + }, + "url": event.change_request.url, + "title": event.change_request.title, + "body": event.change_request.body, + "state": event.change_request.state.value, + "source_branch": event.change_request.source_branch, + "target_branch": event.change_request.target_branch, + "head_sha": event.change_request.head_sha, + "draft": event.change_request.draft, + } + if event.change_request + else None + ), + "comment": (_review_comment_to_dict(event.comment) if event.comment else None), + "review": ( + { + "id": event.review.id, + "state": event.review.state.value, + "body": event.review.body, + "author": event.review.author, + "comments": [_review_comment_to_dict(c) for c in event.review.comments], + } + if event.review + else None + ), + "check": ( + { + "name": event.check.name, + "status": event.check.status.value, + "conclusion": event.check.conclusion.value, + "url": event.check.url, + "logs_url": event.check.logs_url, + "output": event.check.output, + } + if event.check + else None + ), + "check_suite_status": ( + event.check_suite_status.value if event.check_suite_status else None + ), + "raw": event.raw, + } + + +def _review_comment_to_dict(comment: ReviewComment) -> dict[str, Any]: + """Serialize a ReviewComment for queue transport.""" + return { + "id": comment.id, + "body": comment.body, + "author": comment.author, + "path": comment.path, + "line": comment.line, + "resolved": comment.resolved, + "in_reply_to": comment.in_reply_to, + } + + +def _review_comment_from_dict(data: dict[str, Any]) -> ReviewComment: + """Deserialize a ReviewComment from queue transport.""" + return ReviewComment( + id=data["id"], + body=data["body"], + author=data["author"], + path=data.get("path"), + line=data.get("line"), + resolved=data.get("resolved", False), + in_reply_to=data.get("in_reply_to"), + ) + + +def normalized_event_from_dict(data: dict[str, Any]) -> NormalizedEvent: + """Deserialize a NormalizedEvent from queue transport.""" + repo_ref_data = data["repo_ref"] + repo_ref = RepositoryRef( + id=repo_ref_data["id"], + provider=Provider(repo_ref_data["provider"]), + connection=repo_ref_data["connection"], + namespace=repo_ref_data["namespace"], + default_branch=repo_ref_data["default_branch"], + change_request_mode=repo_ref_data["change_request_mode"], + ) + actor = Actor(login=data["actor"]["login"], is_bot=data["actor"]["is_bot"]) + + change_request = None + if data.get("change_request") is not None: + cr = data["change_request"] + change_request = ChangeRequest( + identity=ChangeRequestIdentity( + connection=cr["identity"]["connection"], + repository_id=cr["identity"]["repository_id"], + native_id=cr["identity"]["native_id"], + ), + url=cr["url"], + title=cr["title"], + body=cr["body"], + state=ChangeRequestState(cr["state"]), + source_branch=cr["source_branch"], + target_branch=cr["target_branch"], + head_sha=cr.get("head_sha", ""), + draft=cr["draft"], + ) + + comment = None + if data.get("comment") is not None: + comment = _review_comment_from_dict(data["comment"]) + + review = None + if data.get("review") is not None: + r = data["review"] + review = Review( + id=r["id"], + state=ReviewState(r["state"]), + body=r["body"], + author=r["author"], + comments=[_review_comment_from_dict(c) for c in r.get("comments", [])], + ) + + check = None + if data.get("check") is not None: + ck = data["check"] + check = CheckRun( + name=ck["name"], + status=CheckStatus(ck["status"]), + conclusion=CheckConclusion(ck["conclusion"]), + url=ck.get("url"), + logs_url=ck.get("logs_url"), + output=ck.get("output", {}), + ) + + check_suite_status_value = data.get("check_suite_status") + + return NormalizedEvent( + id=data["id"], + kind=EventKind(data["kind"]), + repo_ref=repo_ref, + actor=actor, + received_at=datetime.fromisoformat(data["received_at"]), + change_request=change_request, + comment=comment, + review=review, + check=check, + check_suite_status=( + CheckStatus(check_suite_status_value) if check_suite_status_value else None + ), + raw=data.get("raw", {}), + ) diff --git a/src/forge/queue/producer.py b/src/forge/queue/producer.py index 55894679f..49129143a 100644 --- a/src/forge/queue/producer.py +++ b/src/forge/queue/producer.py @@ -5,6 +5,7 @@ import redis.asyncio as redis +from forge.integrations.source_control.contracts import NormalizedEvent from forge.models.events import EventSource from forge.orchestrator.checkpointer import get_redis_client from forge.queue.deduplication import DEDUP_KEY_PREFIX, DEDUP_TTL_SECONDS @@ -14,7 +15,13 @@ # Stream names for different event sources JIRA_STREAM = "forge:events:jira" -GITHUB_STREAM = "forge:events:github" +SOURCE_CONTROL_STREAM = "forge:events:source_control" + +# Pre-rename stream name (source-control events used to publish here, and to +# EventSource value "github"). New events never publish to this stream, but +# it may still hold unconsumed entries from before the rename, so the +# consumer keeps draining it -- see queue/consumer.py. +LEGACY_SOURCE_CONTROL_STREAM = "forge:events:github" _PUBLISH_ONCE_SCRIPT = """ local reserved = redis.call('SET', KEYS[1], '1', 'EX', ARGV[1], 'NX') @@ -45,7 +52,7 @@ async def _get_redis(self) -> redis.Redis: def _get_stream_name(self, source: EventSource) -> str: """Get the appropriate stream name for an event source.""" - return JIRA_STREAM if source == EventSource.JIRA else GITHUB_STREAM + return JIRA_STREAM if source == EventSource.JIRA else SOURCE_CONTROL_STREAM async def publish( self, @@ -123,6 +130,54 @@ async def publish_once( logger.info("Published new event %s to %s as %s", event_id, stream, message_id) return str(message_id) + async def publish_event(self, event: NormalizedEvent, ticket_key: str) -> str | None: + """Atomically publish a NormalizedEvent to the source-control stream, + unless its id was already seen. + + GitHub redelivers webhooks on timeouts, 5xx responses, and manual + "Redeliver," so this needs the same SET-NX-then-XADD dedup guarantee + publish_once gives Jira events -- a plain XADD here would silently + reprocess every retried delivery (duplicate PR comments, duplicate + CI-fix attempts, etc). + + Args: + event: The normalized webhook event. + ticket_key: Jira ticket key this event resolves to (extracted by the + caller before publishing). + + Returns: + The Redis stream message ID, or None if event.id was a duplicate. + """ + from forge.queue.models import normalized_event_to_dict # avoid a cycle at import time + + redis_client = await self._get_redis() + message = QueueMessage( + message_id="", + event_id=event.id, + source=EventSource.SOURCE_CONTROL, + event_type=event.kind.value, + ticket_key=ticket_key, + payload=event.raw, + normalized_event=normalized_event_to_dict(event), + ) + fields = message.to_dict() + field_values = [item for pair in fields.items() for item in pair] + message_id = await redis_client.eval( + _PUBLISH_ONCE_SCRIPT, + 2, + f"{DEDUP_KEY_PREFIX}{event.id}", + SOURCE_CONTROL_STREAM, + DEDUP_TTL_SECONDS, + *field_values, + ) + if message_id is None: + logger.info("Skipped duplicate event %s for %s", event.id, SOURCE_CONTROL_STREAM) + return None + logger.info( + "Published new event %s to %s as %s", event.id, SOURCE_CONTROL_STREAM, message_id + ) + return str(message_id) + async def republish(self, message: QueueMessage) -> str: """Republish a message (e.g., for retry). diff --git a/src/forge/workflow/nodes/ci_evaluator.py b/src/forge/workflow/nodes/ci_evaluator.py index f906f48b5..48c393fbf 100644 --- a/src/forge/workflow/nodes/ci_evaluator.py +++ b/src/forge/workflow/nodes/ci_evaluator.py @@ -9,8 +9,15 @@ from forge.api.routes.metrics import record_ci_fix_attempt from forge.config import get_settings -from forge.integrations.github.client import GitHubClient from forge.integrations.jira.client import JiraClient +from forge.integrations.source_control.contracts import ( + CheckConclusion, + CheckRun, + CheckStatus, + RepositoryRef, + SourceControlProvider, +) +from forge.integrations.source_control.errors import NotFoundError from forge.models.workflow import ForgeLabel from forge.prompts import load_prompt from forge.sandbox import ContainerRunner @@ -18,11 +25,13 @@ from forge.workflow.nodes.code_review import run_post_change_review, sync_pr_description from forge.workflow.nodes.error_handler import notify_error from forge.workflow.nodes.workspace_setup import prepare_workspace +from forge.workflow.pr_state import find_active_pull_request from forge.workflow.utils import merge_review_exhaustion, update_state_timestamp from forge.workflow.utils.jira_status import ( post_status_comment, set_review_pending_label, ) +from forge.workflow.utils.source_control import get_adapter, identity_for from forge.workspace.git_ops import GitOperations from forge.workspace.handoff import capture_handoff from forge.workspace.manager import Workspace @@ -47,7 +56,12 @@ async def evaluate_ci_status(state: WorkflowState) -> WorkflowState: ticket_key = state["ticket_key"] current_pr_url = state.get("current_pr_url") pull_requests = state.get("pull_requests", {}) - active_pr = pull_requests.get(state.get("current_repo", "")) + _, active_pr = find_active_pull_request( + pull_requests, + state.get("current_repo", ""), + state.get("current_pr_number"), + current_pr_url, + ) if pull_requests: if not ( isinstance(active_pr, dict) @@ -61,6 +75,7 @@ async def evaluate_ci_status(state: WorkflowState) -> WorkflowState: "ci_status": "failed", "current_node": "ci_evaluator", "last_error": "Active pull request state is inconsistent", + "pending_ci_event": False, } ) pr_urls = [current_pr_url] @@ -83,17 +98,13 @@ async def evaluate_ci_status(state: WorkflowState) -> WorkflowState: logger.info(f"Evaluating CI status for {ticket_key}") - github = GitHubClient() - try: # Checks whose name contains any of these substrings are treated as passing ci_skipped_checks = state.get("ci_skipped_checks", []) - def _is_skipped(check: dict) -> bool: - name = check.get("name", "") - return any(skip.lower() in name.lower() for skip in ci_skipped_checks) + def _is_skipped(check: CheckRun) -> bool: + return any(skip.lower() in check.name.lower() for skip in ci_skipped_checks) - # Check each PR's CI status. # Only *completed* non-skipped checks count toward pass/fail. # Pending checks (e.g. tide, which waits for merge labels) are ignored # once at least one real check has completed — they would block forever. @@ -105,72 +116,62 @@ def _is_skipped(check: dict) -> bool: all_passed = True _any_skipped = False any_still_running = False - failed_checks = [] + failed_checks: list[dict[str, Any]] = [] - for pr_url in pr_urls: - # Parse PR URL to get owner/repo/number - parts = pr_url.rstrip("/").split("/") - owner, repo = parts[-4], parts[-3] - pr_number = int(parts[-1]) + repo_ref, adapter = get_adapter(state.get("current_repo", "")) + identity = identity_for(repo_ref, state.get("current_pr_number")) + change_request = await adapter.get_change_request(repo_ref, identity) + checks = await adapter.get_checks( + repo_ref, change_request.head_sha or change_request.source_branch + ) - # Get PR details for head SHA - pr_data = await github.get_pull_request(owner, repo, pr_number) - head_sha = pr_data.get("head", {}).get("sha", "") + # If no checks exist yet, CI is still pending + if not checks: + logger.info(f"No CI checks registered yet for {ticket_key}, waiting for webhook") + return update_state_timestamp( + { + **state, + "ci_status": "pending", + "current_node": "ci_evaluator", # Stay here, wait for webhook + "pending_ci_event": False, + } + ) - if not head_sha: + for check in checks: + if _is_skipped(check): + logger.info(f"CI check skipped by human override: {check.name}") + _any_skipped = True continue - # Get check runs for the commit - check_runs = await github.get_check_runs(owner, repo, head_sha) + # Ignore permanently-pending meta-checks (e.g. tide) — they + # wait for merge labels, not for CI to pass, and would block + # evaluation indefinitely. + if check.status != CheckStatus.COMPLETED and any( + p in check.name.lower() for p in _permanent_pending + ): + logger.info(f"Ignoring permanently-pending check: {check.name}") + continue - # If no check runs exist yet, CI is still pending - if not check_runs: - logger.info(f"No CI checks registered yet for {pr_url}, waiting for webhook") - return update_state_timestamp( + if check.status != CheckStatus.COMPLETED: + # Real CI check still running — wait for it + all_passed = False + any_still_running = True + logger.info(f"CI still running for {ticket_key}") + elif check.conclusion not in ( + CheckConclusion.SUCCESS, + CheckConclusion.SKIPPED, + CheckConclusion.NEUTRAL, + ): + all_passed = False + failed_checks.append( { - **state, - "ci_status": "pending", - "current_node": "ci_evaluator", # Stay here, wait for webhook - "pending_ci_event": False, + "name": check.name, + "conclusion": check.conclusion.value, + "output": check.output, + "logs_ref": check.logs_url, } ) - for check in check_runs: - if _is_skipped(check): - logger.info(f"CI check skipped by human override: {check.get('name')}") - _any_skipped = True - continue - - check_name = check.get("name", "") - status = check.get("status") - conclusion = check.get("conclusion") - - # Ignore permanently-pending meta-checks (e.g. tide) — they - # wait for merge labels, not for CI to pass, and would block - # evaluation indefinitely. - if status != "completed" and any( - p in check_name.lower() for p in _permanent_pending - ): - logger.info(f"Ignoring permanently-pending check: {check_name}") - continue - - if status != "completed": - # Real CI check still running — wait for it - all_passed = False - any_still_running = True - logger.info(f"CI still running for {pr_url}") - elif conclusion not in ("success", "skipped", "neutral"): - all_passed = False - failed_checks.append( - { - "pr_url": pr_url, - "name": check_name, - "conclusion": conclusion, - "output": check.get("output", {}), - "log_url": check.get("html_url", ""), - } - ) - if all_passed: logger.info(f"All CI checks passed for {ticket_key}") jira = JiraClient() @@ -255,8 +256,6 @@ def _is_skipped(check: dict) -> bool: "retry_count": state.get("retry_count", 0) + 1, "pending_ci_event": False, } - finally: - await github.close() async def attempt_ci_fix(state: WorkflowState) -> WorkflowState: @@ -278,11 +277,18 @@ async def attempt_ci_fix(state: WorkflowState) -> WorkflowState: failed_checks = state.get("ci_failed_checks", []) if not failed_checks: + # Defensive: attempt_ci_fix is only routed to when ci_failed_checks is + # non-empty. If it arrived here empty anyway (e.g. a concurrent/stale + # state update cleared it before this node ran), re-verify against + # live CI rather than asserting the checks passed. + logger.warning( + f"attempt_ci_fix entered with no ci_failed_checks for {ticket_key} " + "— re-verifying CI status instead of assuming pass" + ) return update_state_timestamp( { **state, - "ci_status": "passed", - "current_node": "human_review_gate", + "current_node": "ci_evaluator", "pending_ci_event": False, } ) @@ -304,7 +310,7 @@ async def attempt_ci_fix(state: WorkflowState) -> WorkflowState: fork_owner = state.get("fork_owner", "") fork_repo = state.get("fork_repo", "") try: - workspace_path, _ = prepare_workspace(state) + workspace_path, _ = await prepare_workspace(state) state = {**state, "workspace_path": workspace_path} except Exception as _setup_err: logger.error(f"Workspace setup failed for {ticket_key}: {_setup_err}") @@ -322,11 +328,8 @@ async def attempt_ci_fix(state: WorkflowState) -> WorkflowState: logs_dir = Path(workspace_path) / ".forge" / "logs" failures_file.parent.mkdir(parents=True, exist_ok=True) - github = GitHubClient() - try: - await _fetch_ci_logs_and_artifacts(failed_checks, logs_dir, github) - finally: - await github.close() + repo_ref, adapter = get_adapter(state.get("current_repo", "")) + await _fetch_ci_logs_and_artifacts(failed_checks, logs_dir, repo_ref, adapter) failures_file.write_text(_collect_error_info(failed_checks)) # ── Phase 0: Attribution ───────────────────────────────────────────── @@ -462,7 +465,7 @@ async def attempt_ci_fix(state: WorkflowState) -> WorkflowState: branch_name=state.get("context", {}).get("branch_name", ""), ticket_key=ticket_key, ) - git = GitOperations(workspace) + git = GitOperations(workspace, await adapter.get_git_credentials(repo_ref)) branch_name = state.get("context", {}).get("branch_name", "") if git.has_uncommitted_changes(): @@ -501,13 +504,10 @@ async def attempt_ci_fix(state: WorkflowState) -> WorkflowState: logger.info(f"CI fix pushed for {ticket_key} (attempt {ci_fix_attempt})") record_ci_fix_attempt(repo=state.get("current_repo", "unknown"), result="pushed") - _repo = state.get("current_repo", "/") - _owner, _repo_name = (_repo.split("/") + [""])[:2] await sync_pr_description( state, git, - owner=_owner, - repo=_repo_name, + current_repo=state.get("current_repo", ""), pr_number=state.get("current_pr_number"), attempt=ci_fix_attempt, ) @@ -620,69 +620,59 @@ async def escalate_to_blocked(state: WorkflowState) -> WorkflowState: await jira.close() -def _parse_run_info(log_url: str) -> tuple[str, str, str, str] | None: - """Extract (owner, repo, run_id, job_id) from a GitHub Actions html_url. - - Expects: https://github.com/{owner}/{repo}/actions/runs/{run_id}/job/{job_id} - """ - parts = log_url.rstrip("/").split("/") - if len(parts) < 10 or parts[5] != "actions" or parts[8] != "job": - return None - return parts[3], parts[4], parts[7], parts[9] - - async def _fetch_ci_logs_and_artifacts( failed_checks: list[dict[str, Any]], logs_dir: Path, - github: GitHubClient, + repo_ref: RepositoryRef, + adapter: SourceControlProvider, ) -> None: """Download job logs and run artifacts for all failed checks into logs_dir. - Deduplicates artifact downloads across checks sharing the same run_id. - Silently skips any individual download that fails so one bad log does not - block the whole fix pipeline. + Deduplicates artifact downloads across checks sharing the same run + (``logs_ref``). Silently skips any individual download that fails so one + bad log does not block the whole fix pipeline. """ logs_dir.mkdir(parents=True, exist_ok=True) - fetched_run_ids: set[str] = set() + fetched_run_refs: set[str] = set() - for check in failed_checks: - log_url = check.get("log_url", "") - check_name = check.get("name", "unknown") - info = _parse_run_info(log_url) - if not info: - continue - owner, repo, run_id, job_id = info + for check_record in failed_checks: + check_name = check_record.get("name", "unknown") + logs_ref = check_record.get("logs_ref") safe_name = check_name.replace(" ", "-").replace("/", "-") + check = CheckRun( + name=check_name, + status=CheckStatus.COMPLETED, + conclusion=CheckConclusion.FAILURE, + logs_url=logs_ref, + ) # Job log try: - log_text = await github.get_job_logs(owner, repo, job_id) - (logs_dir / f"{safe_name}-{job_id}.txt").write_text(log_text, errors="replace") + log_text = await adapter.get_check_logs(repo_ref, check) + (logs_dir / f"{safe_name}.txt").write_text(log_text, errors="replace") logger.info(f"Downloaded job log for '{check_name}' ({len(log_text)} chars)") + except NotFoundError: + logger.warning(f"No logs found for '{check_name}'") except Exception as e: logger.warning(f"Could not download job log for '{check_name}': {e}") - # Artifacts (once per run_id) - if run_id in fetched_run_ids: + # Artifacts (once per run) + if not logs_ref or logs_ref in fetched_run_refs: continue - fetched_run_ids.add(run_id) + fetched_run_refs.add(logs_ref) try: - artifacts = await github.get_run_artifacts(owner, repo, run_id) - for artifact in artifacts: - artifact_id = artifact.get("id") - artifact_name = artifact.get("name", str(artifact_id)) + for artifact_name, zip_bytes in await adapter.get_check_artifacts(repo_ref, check): try: - zip_bytes = await github.download_artifact_zip(owner, repo, artifact_id) artifact_dir = logs_dir / artifact_name artifact_dir.mkdir(exist_ok=True) with zipfile.ZipFile(io.BytesIO(zip_bytes)) as zf: zf.extractall(artifact_dir) logger.info(f"Extracted artifact '{artifact_name}' to {artifact_dir}") except Exception as e: - logger.warning(f"Could not download artifact '{artifact_name}': {e}") + logger.warning(f"Could not extract artifact '{artifact_name}': {e}") except Exception as e: - logger.warning(f"Could not list artifacts for run {run_id}: {e}") + logger.warning(f"Could not list artifacts for '{check_name}': {e}") def _collect_error_info(failed_checks: list[dict[str, Any]]) -> str: @@ -697,18 +687,13 @@ def _collect_error_info(failed_checks: list[dict[str, Any]]) -> str: parts = [] for check in failed_checks: + check_name = check.get("name", "unknown") parts.append(f"## {check.get('name', 'Unknown Check')}") parts.append(f"Result: {check.get('conclusion', 'failed')}") - log_url = check.get("log_url", "") - check_name = check.get("name", "unknown") - if log_url: - parts.append(f"Log URL: {log_url}") - info = _parse_run_info(log_url) - if info: - _, _, _, job_id = info - safe_name = check_name.replace(" ", "-").replace("/", "-") - parts.append(f"Log file: `.forge/logs/{safe_name}-{job_id}.txt`") + if check.get("logs_ref"): + safe_name = check_name.replace(" ", "-").replace("/", "-") + parts.append(f"Log file: `.forge/logs/{safe_name}.txt`") output = check.get("output", {}) if output: diff --git a/src/forge/workflow/nodes/code_review.py b/src/forge/workflow/nodes/code_review.py index e4abfc065..778d85416 100644 --- a/src/forge/workflow/nodes/code_review.py +++ b/src/forge/workflow/nodes/code_review.py @@ -12,12 +12,12 @@ from forge.config import get_settings from forge.integrations.agents import ForgeAgent -from forge.integrations.github.client import GitHubClient from forge.integrations.jira.client import JiraClient from forge.prompts import load_prompt from forge.sandbox import ContainerRunner from forge.sandbox.runner import ContainerResult from forge.workflow.utils.jira_status import post_status_comment +from forge.workflow.utils.source_control import get_adapter, identity_for from forge.workspace.git_ops import GitOperations from forge.workspace.manager import Workspace @@ -75,13 +75,15 @@ async def run_post_change_review( skill_name="review-code", ) + repo_ref, adapter = get_adapter(current_repo) git = GitOperations( Workspace( path=Path(workspace_path), repo_name=current_repo, branch_name=branch_name, ticket_key=ticket_key, - ) + ), + await adapter.get_git_credentials(repo_ref), ) if git.has_uncommitted_changes(): @@ -101,8 +103,8 @@ async def run_post_change_review( async def sync_pr_description( state: Any, git: Any, - owner: str, - repo: str, + *, + current_repo: str, pr_number: int | None, attempt: int, ) -> None: @@ -115,8 +117,7 @@ async def sync_pr_description( Args: state: Current workflow state (for ticket_key and audit comment). git: GitOperations instance for the workspace. - owner: Repository owner. - repo: Repository name. + current_repo: Repository identifier (owner/repo or repos.yaml id). pr_number: Pull request number, or None to skip. attempt: Which code-change attempt this follows (0 = initial PR creation). """ @@ -136,11 +137,12 @@ async def sync_pr_description( logger.debug("PR description sync skipped — no commits on branch") return - github = GitHubClient() + repo_ref, adapter = get_adapter(current_repo) + identity = identity_for(repo_ref, pr_number) jira = JiraClient() try: - pr_data = await github.get_pull_request(owner, repo, pr_number) - current_body = pr_data.get("body", "") or "" + change_request = await adapter.get_change_request(repo_ref, identity) + current_body = change_request.body prompt = load_prompt( "sync-pr-description", @@ -153,7 +155,7 @@ async def sync_pr_description( task="sync-pr-description", policy_key="sync_pr_description", prompt=prompt, - context={"owner": owner, "repo": repo, "pr_number": pr_number}, + context={"repo": current_repo, "pr_number": pr_number}, trace_context={ "ticket_key": state.get("ticket_key", ""), "ticket_type": state.get("ticket_type", ""), @@ -171,7 +173,7 @@ async def sync_pr_description( if updated_body: updated_body = agent._strip_preamble(updated_body) if updated_body and updated_body.strip() != current_body.strip(): - await github.update_pull_request(owner, repo, pr_number, body=updated_body) + await adapter.update_change_request(repo_ref, identity, body=updated_body) ticket_key = state.get("ticket_key", "") label = f"CI fix attempt {attempt}" if attempt > 0 else "PR creation" await post_status_comment( @@ -183,7 +185,6 @@ async def sync_pr_description( else: logger.debug(f"PR #{pr_number} description already accurate — no update needed") finally: - await github.close() await jira.close() except Exception as e: diff --git a/src/forge/workflow/nodes/docs_updater.py b/src/forge/workflow/nodes/docs_updater.py index f36efb087..b9cb603c4 100644 --- a/src/forge/workflow/nodes/docs_updater.py +++ b/src/forge/workflow/nodes/docs_updater.py @@ -8,6 +8,7 @@ from forge.sandbox import ContainerRunner from forge.workflow.feature.state import FeatureState as WorkflowState from forge.workflow.utils import merge_review_exhaustion, update_state_timestamp +from forge.workflow.utils.source_control import get_adapter from forge.workspace.git_ops import GitOperations from forge.workspace.manager import Workspace @@ -66,13 +67,15 @@ async def update_documentation(state: WorkflowState) -> WorkflowState: state = merge_review_exhaustion(state, result, ticket_key, "update_docs") + repo_ref, adapter = get_adapter(current_repo) git = GitOperations( Workspace( path=Path(workspace_path), repo_name=current_repo, branch_name=branch_name, ticket_key=ticket_key, - ) + ), + await adapter.get_git_credentials(repo_ref), ) if git.has_uncommitted_changes(): diff --git a/src/forge/workflow/nodes/error_handler.py b/src/forge/workflow/nodes/error_handler.py index 285f2b7ec..782f4d1e7 100644 --- a/src/forge/workflow/nodes/error_handler.py +++ b/src/forge/workflow/nodes/error_handler.py @@ -7,8 +7,10 @@ from typing import Any from forge.integrations.jira.client import JiraClient +from forge.integrations.source_control.errors import SourceControlError from forge.utils.redaction import redact_secrets from forge.workflow.feature.state import FeatureState as WorkflowState +from forge.workflow.utils.source_control import get_adapter, identity_for logger = logging.getLogger(__name__) _MODEL_POLICY_ERROR_PREFIX = "Model policy configuration error:" @@ -101,17 +103,14 @@ async def notify_error( and "/" in current_repo and pr_number ): - from forge.integrations.github.client import GitHubClient - - owner, repo = current_repo.split("/", 1) - github = GitHubClient() try: + repo_ref, adapter = get_adapter(current_repo) + identity = identity_for(repo_ref, int(pr_number)) problem, available = _model_policy_error_parts(error_truncated) fix_command = _model_policy_fix_command(ticket_key, node_name) - await github.create_issue_comment( - owner, - repo, - int(pr_number), + await adapter.create_comment( + repo_ref, + identity, "## 🚨 Action required: Forge model configuration\n\n" "This stage stopped because its project model selection is invalid. " "**The solution is below.**\n\n" @@ -125,13 +124,15 @@ async def notify_error( f"Then add the `forge:retry` label to `{ticket_key}`.", ) logger.info(f"Posted model policy error to {current_repo}#{pr_number}") - except Exception as github_error: + except SourceControlError as sc_error: + logger.warning( + f"Failed to post model policy error to {current_repo}#{pr_number}: {sc_error}" + ) + except Exception as unexpected_error: logger.warning( - f"Failed to post model policy error to {current_repo}#{pr_number}: " - f"{github_error}" + f"Unexpected error posting model policy error to " + f"{current_repo}#{pr_number}: {unexpected_error}" ) - finally: - await github.close() except Exception as e: # Don't fail the workflow if we can't post a comment diff --git a/src/forge/workflow/nodes/git_persistence.py b/src/forge/workflow/nodes/git_persistence.py index 7858e1e70..b7d511823 100644 --- a/src/forge/workflow/nodes/git_persistence.py +++ b/src/forge/workflow/nodes/git_persistence.py @@ -75,13 +75,23 @@ def classify_push_failure(error: Exception) -> PushFailureKind: async def push_to_fork_with_retry( git: GitOperations, *, + use_fork: bool = True, max_attempts: int = 3, initial_delay_seconds: float = 1.0, ) -> None: - """Push a workflow branch, retrying only failures known to be transient.""" + """Push a workflow branch, retrying only failures known to be transient. + + Args: + git: Git operations bound to the workspace. + use_fork: Push to the 'fork' remote (default, fork mode). When False + (direct mode — no fork identity), pushes to 'origin' instead. + """ for attempt in range(1, max_attempts + 1): try: - git.push_to_fork() + if use_fork: + git.push_to_fork() + else: + git.push(force=False, check_conflicts=False) return except Exception as exc: kind = classify_push_failure(exc) @@ -98,6 +108,15 @@ async def push_to_fork_with_retry( await asyncio.sleep(delay) +def use_fork_remote(state: dict) -> bool: + """Whether this workflow's write target is a fork (vs. direct-to-origin). + + Fork identity (fork_owner/fork_repo) is only populated in state for + change_request_mode == "fork" repos; direct-mode repos leave both empty. + """ + return bool(state.get("fork_owner") and state.get("fork_repo")) + + def build_persistence_error_state( state: dict, error: PushPersistenceError, diff --git a/src/forge/workflow/nodes/human_review.py b/src/forge/workflow/nodes/human_review.py index 2be82d97b..194cd6b50 100644 --- a/src/forge/workflow/nodes/human_review.py +++ b/src/forge/workflow/nodes/human_review.py @@ -35,7 +35,12 @@ async def human_review_gate(state: WorkflowState) -> WorkflowState: if state.get("pr_merged"): logger.info(f"PR already merged for {ticket_key}, skipping pause at human_review_gate") return update_state_timestamp( - {**state, "current_node": "human_review_gate", "is_paused": False} + { + **state, + "current_node": "human_review_gate", + "is_paused": False, + "pending_ci_event": False, + } ) updates: dict[str, Any] = {} @@ -44,9 +49,15 @@ async def human_review_gate(state: WorkflowState) -> WorkflowState: try: pr_number = state.get("current_pr_number") if pr_number is not None: + pr_url = state.get("current_pr_url") + if not pr_url: + pr_urls = state.get("pr_urls", []) + pr_url = pr_urls[-1] if pr_urls else None + pr_label = f"Pull request #{pr_number}" + if pr_url: + pr_label = f"[{pr_label}]({pr_url})" message = ( - f"🚀 Pull request #{pr_number} created and submitted. " - "Waiting for CI checks and human review." + f"🚀 {pr_label} created and submitted. Waiting for CI checks and human review." ) else: message = ( @@ -81,6 +92,12 @@ def route_human_review(state: WorkflowState) -> str: Returns: Next node name or END. """ + # Check if merged — takes priority over a stale pending_ci_event so an + # already-merged PR doesn't get routed back through ci_evaluator. + if state.get("pr_merged"): + logger.info(f"PR merged for {state['ticket_key']}") + return "complete_tasks" + # CI webhook arrived while paused at this gate — route through CI cycle if state.get("pending_ci_event"): logger.info(f"CI event pending for {state['ticket_key']}, routing to ci_evaluator") @@ -91,11 +108,6 @@ def route_human_review(state: WorkflowState) -> str: logger.info(f"Changes requested for {state['ticket_key']}") return "implement_review" - # Check if merged - if state.get("pr_merged"): - logger.info(f"PR merged for {state['ticket_key']}") - return "complete_tasks" - # Still waiting for review - END and wait for webhook if state.get("is_paused"): logger.info( diff --git a/src/forge/workflow/nodes/implement_review.py b/src/forge/workflow/nodes/implement_review.py index b447aa039..9a99038b3 100644 --- a/src/forge/workflow/nodes/implement_review.py +++ b/src/forge/workflow/nodes/implement_review.py @@ -8,8 +8,6 @@ from langgraph.graph import END from forge.config import get_settings -from forge.integrations.github.client import GitHubClient -from forge.integrations.github.comment_signature import is_self_comment, resolve_bot_login from forge.integrations.jira.client import JiraClient from forge.prompts import load_prompt from forge.sandbox import ContainerRunner @@ -22,6 +20,7 @@ merge_review_decisions, reply_to_review_decisions, ) +from forge.workflow.utils.source_control import get_adapter, identity_for from forge.workspace.handoff import capture_handoff logger = logging.getLogger(__name__) @@ -86,20 +85,18 @@ def _thread_is_settled( if not comments: return False last_comment = comments[-1] - return is_self_comment( - sender_login=last_comment.get("author", ""), - comment_body=last_comment.get("body", ""), - bot_login=bot_login, - prefix=prefix, + sender_login = last_comment.get("author", "") + comment_body = last_comment.get("body", "") + return sender_login.casefold() == bot_login.casefold() or bool( + prefix and comment_body.startswith(prefix) ) async def _fetch_pr_review_comments( - owner: str, - repo: str, + current_repo: str, pr_number: int, review_body: str, - review_comments: list[dict[str, Any]] | None = None, + review_comments: list[dict[str, Any]] | set[str] | None = None, ) -> str: """Fetch all PR review comments and format them for the analysis container. @@ -109,8 +106,7 @@ async def _fetch_pr_review_comments( excluded so the analysis agent isn't re-fed the same resolved feedback. Args: - owner: Repository owner. - repo: Repository name. + current_repo: Repository identifier (owner/repo or repos.yaml id). pr_number: PR number. review_body: The review summary body from the webhook. review_comments: Prior per-thread decisions from earlier review cycles. @@ -118,35 +114,27 @@ async def _fetch_pr_review_comments( Returns: Formatted markdown string of all review feedback. """ - github = GitHubClient() try: - threads = await github.get_pull_request_review_threads(owner, repo, pr_number) + repo_ref, adapter = get_adapter(current_repo) + identity = identity_for(repo_ref, pr_number) + reviews = await adapter.get_review_thread_comments(repo_ref, identity) except Exception as e: logger.warning(f"Could not fetch inline review comments: {e}") - try: - comments = await github.get_pull_request_review_comments(owner, repo, pr_number) - threads = [ - { - "thread_id": comment.get("thread_id") or f"comment-{comment.get('id')}", - "path": comment.get("path", ""), - "line": comment.get("line") - or comment.get("original_line") - or comment.get("position"), - "comments": [ - { - "comment_id": comment.get("comment_id") or comment.get("id"), - "body": comment.get("body", ""), - "author": (comment.get("user") or {}).get("login", ""), - } - ], - } - for comment in comments - ] - except Exception as fallback_error: - logger.warning("REST review comment fallback failed: %s", fallback_error) - threads = [] - finally: - await github.close() + reviews = [] + + processed_thread_ids = review_comments if isinstance(review_comments, set) else set() + threads: list[dict[str, Any]] = [ + { + "thread_id": review.id, + "path": review.comments[0].path if review.comments else "", + "line": review.comments[0].line if review.comments else None, + "comments": [ + {"comment_id": c.id, "body": c.body, "author": c.author} for c in review.comments + ], + } + for review in reviews + if review.id not in processed_thread_ids + ] lines = ["# PR Review Feedback\n"] @@ -155,10 +143,10 @@ async def _fetch_pr_review_comments( lines.append(review_body.strip()) lines.append("\n") - review_comments = review_comments or [] + prior_decisions = review_comments if isinstance(review_comments, list) else [] dispositions_by_thread = { item["thread_id"]: item.get("disposition") - for item in review_comments + for item in prior_decisions if item.get("thread_id") } @@ -167,13 +155,11 @@ async def _fetch_pr_review_comments( dispositions_by_thread.get(thread["thread_id"]) in ("contest", "clarify") for thread in threads ): - login_client = GitHubClient() try: - bot_login = await resolve_bot_login(login_client) + repo_ref, adapter = get_adapter(current_repo) + bot_login = (await adapter.get_authenticated_identity(repo_ref)).login except Exception as e: logger.warning(f"Could not resolve Forge's bot login: {e}") - finally: - await login_client.close() prefix = get_settings().forge_bot_comment_prefix threads = [ @@ -223,7 +209,7 @@ def _load_review_decisions(workspace_path: str) -> list[dict[str, Any]]: async def _reply_to_review_threads( - *, owner: str, repo: str, pr_number: int | None, decisions: list[dict[str, Any]] + *, current_repo: str, pr_number: int | None, decisions: list[dict[str, Any]] ) -> None: """Post decision responses in their originating GitHub review threads.""" # skip_addressed is defense-in-depth, not the primary guard: decisions here @@ -232,7 +218,7 @@ async def _reply_to_review_threads( # thread out of the analysis input entirely) is what actually prevents # re-posting today. await reply_to_review_decisions( - repo_full_name=f"{owner}/{repo}", + current_repo=current_repo, pr_number=pr_number, decisions=decisions, skip_addressed=True, @@ -273,7 +259,7 @@ async def implement_review(state: WorkflowState) -> WorkflowState: try: try: - workspace_path, git = prepare_workspace(state) + workspace_path, git = await prepare_workspace(state) state = {**state, "workspace_path": workspace_path} except ValueError as e: return update_state_timestamp( @@ -286,17 +272,14 @@ async def implement_review(state: WorkflowState) -> WorkflowState: ) # ── Phase 0: Fetch all PR review comments from GitHub ───────────────── - _owner, _, _repo = current_repo.partition("/") await _post_review_addressing_comment( ticket_key=ticket_key, - owner=_owner, - repo=_repo, + current_repo=current_repo, pr_number=pr_number, ) review_comments_text = await _fetch_pr_review_comments( - owner=_owner, - repo=_repo, + current_repo=current_repo, pr_number=pr_number or 0, review_body=feedback_comment, review_comments=state.get("review_comments", []), @@ -342,14 +325,10 @@ async def implement_review(state: WorkflowState) -> WorkflowState: item for item in decisions if item["disposition"] in ("contest", "clarify") ] await _reply_to_review_threads( - owner=_owner, - repo=_repo, + current_repo=current_repo, pr_number=pr_number, decisions=response_decisions, ) - for decision in contested_comments: - decision["status"] = "addressed" - # Backward-compatible fallback if an older analysis prompt writes only # the legacy objections file. New analysis never blocks accepted work. objections_path = Path(workspace_path) / _REVIEW_OBJECTIONS_FILE @@ -360,8 +339,7 @@ async def implement_review(state: WorkflowState) -> WorkflowState: await _post_review_objection( state=state, objections=objections_text, - owner=_owner, - repo=_repo, + current_repo=current_repo, pr_number=pr_number, ) contested_comments = [{"text": objections_text}] @@ -433,8 +411,7 @@ async def implement_review(state: WorkflowState) -> WorkflowState: await sync_pr_description( state, git, - owner=_owner, - repo=_repo, + current_repo=current_repo, pr_number=pr_number, attempt=0, ) @@ -450,8 +427,7 @@ async def implement_review(state: WorkflowState) -> WorkflowState: if not decision.get("response"): decision["response"] = accepted_response await _reply_to_review_threads( - owner=_owner, - repo=_repo, + current_repo=current_repo, pr_number=pr_number, decisions=accepted_decisions, ) @@ -489,8 +465,7 @@ async def implement_review(state: WorkflowState) -> WorkflowState: async def _post_review_addressing_comment( ticket_key: str, - owner: str, - repo: str, + current_repo: str, pr_number: int | None, ) -> None: """Post a non-triggering PR status update when review work starts.""" @@ -499,11 +474,9 @@ async def _post_review_addressing_comment( return try: - github = GitHubClient() - try: - await github.create_issue_comment(owner, repo, pr_number, _REVIEW_ADDRESSING_COMMENT) - finally: - await github.close() + repo_ref, adapter = get_adapter(current_repo) + identity = identity_for(repo_ref, pr_number) + await adapter.create_comment(repo_ref, identity, _REVIEW_ADDRESSING_COMMENT) except Exception as e: logger.warning(f"Failed to post review addressing PR status for {ticket_key}: {e}") @@ -511,14 +484,12 @@ async def _post_review_addressing_comment( async def _post_review_objection( state: Any, objections: str, - owner: str, - repo: str, + current_repo: str, pr_number: int | None, ) -> None: """Post the agent's review objections to the PR and Jira.""" ticket_key = state.get("ticket_key", "") try: - github = GitHubClient() jira = JiraClient() try: comment = ( @@ -527,7 +498,9 @@ async def _post_review_objection( f"*Please confirm whether to proceed as requested or withdraw.*" ) if pr_number: - await github.create_issue_comment(owner, repo, pr_number, comment) + repo_ref, adapter = get_adapter(current_repo) + identity = identity_for(repo_ref, pr_number) + await adapter.create_comment(repo_ref, identity, comment) await post_status_comment( jira, ticket_key, @@ -535,7 +508,6 @@ async def _post_review_objection( f"Objection posted on PR #{pr_number}. Awaiting confirmation.", ) finally: - await github.close() await jira.close() except Exception as e: logger.warning(f"Failed to post review objection: {e}") diff --git a/src/forge/workflow/nodes/implementation.py b/src/forge/workflow/nodes/implementation.py index 7e5abefbf..e89f824f4 100644 --- a/src/forge/workflow/nodes/implementation.py +++ b/src/forge/workflow/nodes/implementation.py @@ -23,6 +23,7 @@ PushPersistenceError, build_persistence_error_state, push_to_fork_with_retry, + use_fork_remote, ) from forge.workflow.nodes.workspace_setup import prepare_workspace from forge.workflow.utils import merge_review_exhaustion, update_state_timestamp @@ -62,7 +63,7 @@ async def implement_task(state: WorkflowState) -> WorkflowState: try: git: GitOperations - workspace_path, git = prepare_workspace(state) + workspace_path, git = await prepare_workspace(state) state = {**state, "workspace_path": workspace_path} except Exception as exc: logger.error("Unable to prepare implementation workspace for %s: %s", ticket_key, exc) @@ -80,7 +81,7 @@ async def implement_task(state: WorkflowState) -> WorkflowState: ) if state.get("implementation_push_pending") and same_workspace_survived: try: - await push_to_fork_with_retry(git) + await push_to_fork_with_retry(git, use_fork=use_fork_remote(state)) except PushPersistenceError as exc: return update_state_timestamp( build_persistence_error_state(state, exc, retry_node=implementation_node) @@ -141,7 +142,7 @@ async def implement_task(state: WorkflowState) -> WorkflowState: ) git.stage_all() git.commit(f"[{ticket_key}] chore: commit uncommitted changes after implementation") - await push_to_fork_with_retry(git) + await push_to_fork_with_retry(git, use_fork=use_fork_remote(state)) except PushPersistenceError as exc: return update_state_timestamp( build_persistence_error_state(state, exc, retry_node=implementation_node) @@ -235,7 +236,7 @@ async def implement_task(state: WorkflowState) -> WorkflowState: # Persist each task commit before checkpointing. A subsequent task # or local review may resume on a worker with a different filesystem. try: - await push_to_fork_with_retry(git) + await push_to_fork_with_retry(git, use_fork=use_fork_remote(state)) except PushPersistenceError as exc: pending_state = { **state, diff --git a/src/forge/workflow/nodes/local_reviewer.py b/src/forge/workflow/nodes/local_reviewer.py index 9f876bf41..42526d158 100644 --- a/src/forge/workflow/nodes/local_reviewer.py +++ b/src/forge/workflow/nodes/local_reviewer.py @@ -13,6 +13,7 @@ PushPersistenceError, build_persistence_error_state, push_to_fork_with_retry, + use_fork_remote, ) from forge.workflow.nodes.review_utils import ( next_review_attempt, @@ -106,7 +107,7 @@ async def local_review_changes(state: WorkflowState) -> WorkflowState: local_workspace_survived = bool(recorded_workspace and Path(recorded_workspace).exists()) try: - workspace_path, git = prepare_workspace(state) + workspace_path, git = await prepare_workspace(state) state = {**state, "workspace_path": workspace_path} except Exception as exc: logger.error("Unable to prepare local-review workspace for %s: %s", ticket_key, exc) @@ -121,7 +122,7 @@ async def local_review_changes(state: WorkflowState) -> WorkflowState: ) if state.get("review_push_pending") and same_workspace_survived: try: - await push_to_fork_with_retry(git) + await push_to_fork_with_retry(git, use_fork=use_fork_remote(state)) except PushPersistenceError as exc: return _review_persistence_error_state(state, exc) updates = state.get("review_push_pending_updates", {}) @@ -417,7 +418,7 @@ async def _persist_review_result( ) -> WorkflowState: """Persist review changes before applying the review's routing decision.""" try: - await push_to_fork_with_retry(git) + await push_to_fork_with_retry(git, use_fork=use_fork_remote(state)) except PushPersistenceError as exc: pending_state = { **state, diff --git a/src/forge/workflow/nodes/pr_creation.py b/src/forge/workflow/nodes/pr_creation.py index f09cd9893..1c1961812 100644 --- a/src/forge/workflow/nodes/pr_creation.py +++ b/src/forge/workflow/nodes/pr_creation.py @@ -2,14 +2,19 @@ import contextlib import logging -from dataclasses import dataclass +from dataclasses import replace from pathlib import Path from typing import Any from forge.config import get_settings from forge.integrations.agents import ForgeAgent -from forge.integrations.github.client import GitHubClient, PullRequestCreationResult from forge.integrations.jira.client import JiraClient +from forge.integrations.source_control.contracts import ( + ChangeRequest, + RepositoryRef, + SourceControlProvider, + WriteTarget, +) from forge.models.workflow import ForgeLabel, TicketType from forge.orchestrator.checkpointer import set_pr_ticket_index from forge.prompts import load_prompt @@ -18,6 +23,7 @@ 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.workflow.utils.source_control import get_adapter, identity_for from forge.workspace.git_ops import GitOperations from forge.workspace.manager import Workspace @@ -26,64 +32,31 @@ logger = logging.getLogger(__name__) -@dataclass(frozen=True) -class PullRequestTarget: - """Resolved upstream and fork repository for PR creation.""" +async def prepare_write_target( + adapter: SourceControlProvider, repo_ref: RepositoryRef, git: GitOperations +) -> WriteTarget: + """Ensure a fork exists/synced and register its remote for local pushes.""" + target = await adapter.ensure_write_target(repo_ref) + if target.fork_owner and target.fork_repo: + git.add_fork_remote(target.fork_owner, target.fork_repo) + return target - owner: str - repo: str - fork_owner: str - fork_repo: str - -async def prepare_pull_request_target( - github: GitHubClient, - git: GitOperations, - current_repo: str, -) -> PullRequestTarget: - """Prepare a fork remote for opening a pull request from the current workspace.""" - if not current_repo or "/" not in current_repo: - raise ValueError( - f"Invalid repository format '{current_repo}': must be in owner/repo format" - ) - - owner, repo = current_repo.split("/", 1) - - logger.info(f"Getting or creating fork for {current_repo}") - fork_data = await github.get_or_create_fork(owner, repo) - fork_owner = fork_data["owner"]["login"] - fork_repo = fork_data["name"] - - await github.sync_fork_with_upstream(fork_owner, fork_repo) - git.add_fork_remote(fork_owner, fork_repo) - - return PullRequestTarget( - owner=owner, - repo=repo, - fork_owner=fork_owner, - fork_repo=fork_repo, - ) - - -async def open_pull_request_from_fork( - github: GitHubClient, - target: PullRequestTarget, +async def open_change_request( + adapter: SourceControlProvider, + repo_ref: RepositoryRef, + target: WriteTarget, *, branch_name: str, title: str, body: str, base: str = "main", draft: bool = False, -) -> PullRequestCreationResult: - """Open a pull request from the prepared fork branch to upstream.""" - return await github.create_pull_request( - owner=target.owner, - repo=target.repo, - title=title, - body=body, - head=f"{target.fork_owner}:{branch_name}", - base=base, - draft=draft, +) -> ChangeRequest: + """Open the PR from the pushed workspace branch.""" + target = replace(target, head_ref=branch_name, base_branch=base) + return await adapter.create_change_request( + repo_ref=repo_ref, target=target, title=title, body=body, draft=draft ) @@ -171,10 +144,11 @@ async def create_pull_request(state: WorkflowState) -> WorkflowState: logger.info(f"Creating PR for {ticket_key} ({len(implemented_tasks)} tasks)") - github = GitHubClient() jira = JiraClient() try: + repo_ref, adapter = get_adapter(current_repo) + # Set up workspace reference context = state.get("context", {}) branch_name = context.get("branch_name", "") @@ -185,9 +159,9 @@ async def create_pull_request(state: WorkflowState) -> WorkflowState: branch_name=branch_name, ticket_key=ticket_key, ) - git = GitOperations(workspace) + git = GitOperations(workspace, await adapter.get_git_credentials(repo_ref)) - pr_target = await prepare_pull_request_target(github, git, current_repo) + target = await prepare_write_target(adapter, repo_ref, git) # Check for merge conflicts before pushing has_conflicts, conflicting_files = await check_merge_conflicts(git, default_branch) @@ -216,8 +190,12 @@ async def create_pull_request(state: WorkflowState) -> WorkflowState: } ) - # Push branch to fork (not origin) - git.push_to_fork() + # Push branch to the write target: the fork in fork mode, origin in + # direct mode (direct-mode repos have no "fork" remote to push to). + if target.fork_owner and target.fork_repo: + git.push_to_fork() + else: + git.push(force=False) # Build PR title — fetch live summary from Jira as source of truth ticket_summary = "" @@ -236,9 +214,10 @@ async def create_pull_request(state: WorkflowState) -> WorkflowState: project_key = ticket_key.split("-")[0] if "-" in ticket_key else ticket_key is_draft = await jira.is_repo_draft(project_key, current_repo) - pr_result = await open_pull_request_from_fork( - github, - pr_target, + change_request = await open_change_request( + adapter, + repo_ref, + target, branch_name=branch_name, title=pr_title, body=pr_body, @@ -246,8 +225,9 @@ async def create_pull_request(state: WorkflowState) -> WorkflowState: draft=is_draft, ) - pr_url = pr_result.pr.get("html_url", "") - pr_number = pr_result.pr.get("number") + pr_url = change_request.url + native_id = change_request.identity.native_id + pr_number = int(native_id) if native_id is not None else None # Log PR number extraction status if pr_number is not None: @@ -276,8 +256,8 @@ async def create_pull_request(state: WorkflowState) -> WorkflowState: if pr_number is not None: logger.info(f"Created PR #{pr_number}: {pr_url}") - if pr_result.created: - await _post_pr_commands_comment(github, pr_target, pr_number) + if change_request.created: + await _post_pr_commands_comment(adapter, repo_ref, pr_number) else: logger.info(f"Created PR (number unavailable): {pr_url}") @@ -300,8 +280,7 @@ async def create_pull_request(state: WorkflowState) -> WorkflowState: await sync_pr_description( state, git, - owner=pr_target.owner, - repo=pr_target.repo, + current_repo=current_repo, pr_number=pr_number, attempt=0, ) @@ -314,17 +293,15 @@ async def create_pull_request(state: WorkflowState) -> WorkflowState: ) if exhaustion_section: try: - pr_data = await github.get_pull_request( - pr_target.owner, pr_target.repo, pr_number - ) - current_body = pr_data.get("body", "") or "" + identity = identity_for(repo_ref, pr_number) + current_pr = await adapter.get_change_request(repo_ref, identity) + current_body = current_pr.body if "## Auto-Review Notes" in current_body: logger.debug("PR body already contains Auto-Review Notes — skipping append") else: - await github.update_pull_request( - pr_target.owner, - pr_target.repo, - pr_number, + await adapter.update_change_request( + repo_ref, + identity, body=current_body + "\n\n" + exhaustion_section, ) logger.info("Appended auto-review exhaustion section to PR body") @@ -338,8 +315,8 @@ async def create_pull_request(state: WorkflowState) -> WorkflowState: "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, + "fork_owner": target.fork_owner, + "fork_repo": target.fork_repo, "current_node": "teardown_workspace", "last_error": None, } @@ -355,13 +332,12 @@ async def create_pull_request(state: WorkflowState) -> WorkflowState: "retry_count": state.get("retry_count", 0) + 1, } finally: - await github.close() await jira.close() async def _post_pr_commands_comment( - github: GitHubClient, - pr_target: PullRequestTarget, + adapter: SourceControlProvider, + repo_ref: RepositoryRef, pr_number: int, ) -> None: """Post informational PR commands comment on a newly created pull request.""" @@ -374,12 +350,7 @@ async def _post_pr_commands_comment( "* `/forge unskip-gate ` - Remove a previously set CI check skip.\n\n" "Feel free to use these commands to manage your workflow!" ) - await github.create_issue_comment( - owner=pr_target.owner, - repo=pr_target.repo, - issue_number=pr_number, - body=comment_body, - ) + await adapter.create_comment(repo_ref, identity_for(repo_ref, pr_number), comment_body) logger.info(f"Posted informational command comment on newly created PR #{pr_number}") except Exception as e: logger.warning(f"Failed to post informational command comment on PR #{pr_number}: {e}") diff --git a/src/forge/workflow/nodes/proposal_pr.py b/src/forge/workflow/nodes/proposal_pr.py index c3c9bb31c..647608bfe 100644 --- a/src/forge/workflow/nodes/proposal_pr.py +++ b/src/forge/workflow/nodes/proposal_pr.py @@ -1,15 +1,15 @@ """Shared publication helpers for PRD and specification proposal PRs.""" -import hashlib import logging -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import Any -from forge.integrations.github.client import GitHubClient from forge.integrations.jira.client import JiraClient, pr_interaction_options +from forge.integrations.source_control.errors import NotFoundError from forge.models.workflow import ForgeLabel from forge.orchestrator.checkpointer import set_pr_ticket_index from forge.workflow.utils.jira_status import post_status_comment +from forge.workflow.utils.source_control import get_adapter, identity_for logger = logging.getLogger(__name__) @@ -58,35 +58,34 @@ async def create_proposal_pr( proposals_path: str, ) -> dict[str, Any]: """Publish an artifact branch to a fork and open its upstream PR.""" - upstream_owner, upstream_repo = proposals_repo.split("/", 1) branch = f"forge/{artifact.branch_segment}/{ticket_key.lower()}" file_path = "/".join(filter(None, [proposals_path, ticket_key, artifact.file_name])) - gh = GitHubClient() + repo_ref, adapter = get_adapter(proposals_repo) jira = JiraClient() try: - fork = await gh.get_or_create_fork(upstream_owner, upstream_repo) - fork_owner = fork["owner"]["login"] - fork_repo = fork["name"] - upstream = await gh.get_repository(upstream_owner, upstream_repo) - default_branch = upstream.get("default_branch") or "main" - synced = await gh.sync_fork_with_upstream(fork_owner, fork_repo, branch=default_branch) - if not synced: - raise RuntimeError( - f"Could not synchronize proposal fork {fork_owner}/{fork_repo} " - f"(branch {default_branch}) with upstream" + default_branch = await adapter.resolve_default_branch(repo_ref) + target = await adapter.ensure_write_target(repo_ref) + fork_owner = target.fork_owner or "" + fork_repo = target.fork_repo or "" + fork_ref = ( + replace( + repo_ref, + id=f"{fork_owner}/{fork_repo}", + namespace=f"{fork_owner}/{fork_repo}", + change_request_mode="direct", ) + if fork_owner and fork_repo + else repo_ref + ) - await gh.create_branch(fork_owner, fork_repo, branch, base=default_branch) - existing_file = await gh.get_file_contents(fork_owner, fork_repo, file_path, branch) - await gh.create_or_update_file( - owner=fork_owner, - repo=fork_repo, - path=file_path, - content=content, - message=f"Add {artifact.title_name} for {ticket_key}", - branch=branch, - sha=existing_file["sha"] if existing_file else None, + await adapter.create_branch(fork_ref, branch, default_branch) + await adapter.put_file( + fork_ref, + file_path, + content, + f"Add {artifact.title_name} for {ticket_key}", + branch, ) pr_body = ( f"**{artifact.title_name} for [{ticket_key}]" @@ -97,17 +96,16 @@ async def create_proposal_pr( "Leave comments on this PR to provide feedback — " f"Forge will regenerate the {artifact.title_name} and push updated commits." ) - pr_result = await gh.create_pull_request( - owner=upstream_owner, - repo=upstream_repo, + change_request = await adapter.create_change_request( + repo_ref, + replace(target, head_ref=branch, base_branch=default_branch), title=f"[{ticket_key}] {artifact.title_name}: {summary}", body=pr_body, - head=f"{fork_owner}:{branch}", - base=default_branch, ) - pr_url = pr_result.pr["html_url"] - pr_number = pr_result.pr["number"] + pr_url = change_request.url + native_id = change_request.identity.native_id + pr_number = int(native_id) if native_id is not None else None await set_pr_ticket_index(pr_url, ticket_key) await jira.set_workflow_label(ticket_key, artifact.pending_label) await post_status_comment( @@ -121,24 +119,20 @@ async def create_proposal_pr( return { f"{prefix}_pr_url": pr_url, f"{prefix}_pr_number": pr_number, - f"{prefix}_pr_repo": proposals_repo, + # Canonical namespace, not the raw (possibly repos.yaml-alias) + # proposals_repo -- webhook matching (worker._is_prd_pr_event / + # _is_spec_pr_event) compares this against event.repo_ref.namespace, + # which is always canonical. + f"{prefix}_pr_repo": repo_ref.namespace, f"{prefix}_pr_fork_owner": fork_owner, f"{prefix}_pr_fork_repo": fork_repo, f"{prefix}_pr_branch": branch, f"{prefix}_pr_file_path": file_path, } finally: - await gh.close() await jira.close() -def _git_blob_sha(content: str) -> str: - """Compute the git blob SHA-1 for a string, matching how Git stores blobs.""" - content_bytes = content.encode() - header = f"blob {len(content_bytes)}\0".encode() - return hashlib.sha1(header + content_bytes).hexdigest() - - async def update_proposal_pr( *, artifact: ProposalArtifact, @@ -152,58 +146,64 @@ async def update_proposal_pr( True if the file was updated, False if the content was unchanged. """ prefix = artifact.state_prefix - upstream_owner, upstream_repo = state[f"{prefix}_pr_repo"].split("/", 1) - owner = state.get(f"{prefix}_pr_fork_owner") or upstream_owner - repo = state.get(f"{prefix}_pr_fork_repo") or upstream_repo + upstream_repo = state[f"{prefix}_pr_repo"] + fork_owner = state.get(f"{prefix}_pr_fork_owner") + fork_repo = state.get(f"{prefix}_pr_fork_repo") branch = state[f"{prefix}_pr_branch"] pr_number = state[f"{prefix}_pr_number"] file_path = state[f"{prefix}_pr_file_path"] - gh = GitHubClient() + repo_ref, adapter = get_adapter(upstream_repo) + write_ref = ( + replace( + repo_ref, + id=f"{fork_owner}/{fork_repo}", + namespace=f"{fork_owner}/{fork_repo}", + change_request_mode="direct", + ) + if fork_owner and fork_repo + else repo_ref + ) + identity = identity_for(repo_ref, pr_number) + try: - file_meta = await gh.get_file_contents(owner, repo, file_path, branch) - if not file_meta: - logger.warning( - "Could not find %s file %s on branch %s", - artifact.title_name, - file_path, - branch, - ) - return False + existing_content = await adapter.get_file(write_ref, file_path, branch) + except NotFoundError: + logger.warning( + "Could not find %s file %s on branch %s", + artifact.title_name, + file_path, + branch, + ) + return False - if _git_blob_sha(content) == file_meta["sha"]: - logger.warning( - "Regenerated %s for %s is unchanged — skipping commit", - artifact.title_name, - ticket_key, - ) - await gh.create_issue_comment( - upstream_owner, - upstream_repo, - pr_number, - f"Forge reviewed the feedback but the regenerated " - f"{artifact.published_name} was unchanged. The feedback may " - f"require manual revision, or it may have already been " - f"addressed in a previous revision.", - ) - return False - - await gh.create_or_update_file( - owner=owner, - repo=repo, - path=file_path, - content=content, - message=(f"Revise {artifact.title_name} for {ticket_key} based on feedback"), - branch=branch, - sha=file_meta["sha"], + if existing_content == content: + logger.warning( + "Regenerated %s for %s is unchanged — skipping commit", + artifact.title_name, + ticket_key, ) - await gh.create_issue_comment( - upstream_owner, - upstream_repo, - pr_number, - f"{artifact.published_name} has been revised based on feedback. " - "Please review the updated version.", + await adapter.create_comment( + repo_ref, + identity, + f"Forge reviewed the feedback but the regenerated " + f"{artifact.published_name} was unchanged. The feedback may " + f"require manual revision, or it may have already been " + f"addressed in a previous revision.", ) - return True - finally: - await gh.close() + return False + + await adapter.put_file( + write_ref, + file_path, + content, + f"Revise {artifact.title_name} for {ticket_key} based on feedback", + branch, + ) + await adapter.create_comment( + repo_ref, + identity, + f"{artifact.published_name} has been revised based on feedback. " + "Please review the updated version.", + ) + return True diff --git a/src/forge/workflow/nodes/qa_handler.py b/src/forge/workflow/nodes/qa_handler.py index 183ef287e..f60eb759f 100644 --- a/src/forge/workflow/nodes/qa_handler.py +++ b/src/forge/workflow/nodes/qa_handler.py @@ -5,10 +5,10 @@ from datetime import UTC, datetime from forge.integrations.agents import ForgeAgent -from forge.integrations.github.client import GitHubClient from forge.integrations.jira.client import JiraClient from forge.workflow.feature.state import FeatureState as WorkflowState from forge.workflow.utils import update_state_timestamp +from forge.workflow.utils.source_control import get_adapter, identity_for logger = logging.getLogger(__name__) @@ -32,12 +32,8 @@ async def _post_qa_response( pr_target = _artifact_pr_target(state, artifact_type) if pr_target: repo_full, pr_number = pr_target - owner, repo_name = repo_full.split("/", 1) - gh = GitHubClient() - try: - await gh.create_issue_comment(owner, repo_name, pr_number, body) - finally: - await gh.close() + repo_ref, adapter = get_adapter(repo_full) + await adapter.create_comment(repo_ref, identity_for(repo_ref, pr_number), body) else: await jira.add_comment(ticket_key, body) diff --git a/src/forge/workflow/nodes/rebase.py b/src/forge/workflow/nodes/rebase.py index 426cc8c22..8695c7e4f 100644 --- a/src/forge/workflow/nodes/rebase.py +++ b/src/forge/workflow/nodes/rebase.py @@ -10,8 +10,12 @@ import logging from forge.config import get_settings -from forge.integrations.github.client import GitHubClient from forge.integrations.jira.client import JiraClient +from forge.integrations.source_control.contracts import ( + ChangeRequestIdentity, + RepositoryRef, + SourceControlProvider, +) from forge.prompts import load_prompt from forge.sandbox import ContainerRunner from forge.workflow.feature.state import FeatureState as WorkflowState @@ -21,11 +25,20 @@ ) from forge.workflow.utils import merge_review_exhaustion, update_state_timestamp from forge.workflow.utils.jira_status import post_status_comment +from forge.workflow.utils.source_control import get_adapter, identity_for from forge.workspace.git_ops import GitOperations logger = logging.getLogger(__name__) +async def _fetch_pr_body( + adapter: SourceControlProvider, repo_ref: RepositoryRef, identity: ChangeRequestIdentity +) -> str: + """Fetch the current PR description, used as context for conflict resolution.""" + change_request = await adapter.get_change_request(repo_ref, identity) + return change_request.body + + async def rebase_pr(state: WorkflowState) -> WorkflowState: """Merge main into the PR branch, resolving conflicts with AI if needed. @@ -42,45 +55,54 @@ async def rebase_pr(state: WorkflowState) -> WorkflowState: pr_number = state.get("current_pr_number") rebase_return_node = state.get("rebase_return_node", "ci_evaluator") - if not current_repo or not fork_owner or not fork_repo or not pr_number: - logger.error(f"Cannot rebase {ticket_key}: missing PR/fork state") + if not current_repo or not pr_number: + logger.error(f"Cannot rebase {ticket_key}: missing PR state") return update_state_timestamp( { **state, "current_node": rebase_return_node, "rebase_return_node": None, - "last_error": "Cannot rebase: missing PR or fork information in workflow state", + "last_error": "Cannot rebase: missing repo or PR number in workflow state", } ) + use_fork = bool(fork_owner and fork_repo) + push_remote = "fork" if use_fork else "origin" - owner, repo = current_repo.split("/", 1) settings = get_settings() jira = JiraClient() - github = GitHubClient() try: + repo_ref, adapter = get_adapter(current_repo) + identity = identity_for(repo_ref, pr_number) + # Set up workspace: clone, add fork remote, checkout PR branch manager = get_workspace_manager() workspace = manager.create_workspace(repo_name=current_repo, ticket_key=ticket_key) - git = GitOperations(workspace) + git = GitOperations(workspace, await adapter.get_git_credentials(repo_ref)) git.clone() write_workspace_identity( workspace.path, ticket_key=ticket_key, repo_name=current_repo, ) - git.add_fork_remote(fork_owner, fork_repo) + if use_fork: + git.add_fork_remote(fork_owner, fork_repo) - if git.remote_branch_exists(workspace.branch_name, remote="fork"): - git.checkout_branch(workspace.branch_name, remote="fork") + if git.remote_branch_exists(workspace.branch_name, remote=push_remote): + git.checkout_branch(workspace.branch_name, remote=push_remote) else: - logger.error(f"Branch {workspace.branch_name} not found on fork") + logger.error(f"Branch {workspace.branch_name} not found on {push_remote}") return update_state_timestamp( { **state, "current_node": rebase_return_node, "rebase_return_node": None, - "last_error": f"Branch {workspace.branch_name} not found on fork {fork_owner}/{fork_repo}", + "last_error": ( + f"Branch {workspace.branch_name} not found on {push_remote} " + f"{fork_owner}/{fork_repo}" + if use_fork + else f"Branch {workspace.branch_name} not found on {push_remote}" + ), } ) @@ -104,12 +126,14 @@ async def rebase_pr(state: WorkflowState) -> WorkflowState: # Clean merge — push it logger.info(f"{ticket_key}: clean merge with main, pushing") - git.push_to_fork(force=True) + if use_fork: + git.push_to_fork(force=True) + else: + git.push(force=True, check_conflicts=False) - await github.create_issue_comment( - owner, - repo, - pr_number, + await adapter.create_comment( + repo_ref, + identity, "Branch has been rebased onto main (no conflicts). CI should re-run.", ) await post_status_comment( @@ -139,10 +163,12 @@ async def rebase_pr(state: WorkflowState) -> WorkflowState: pr_description = "" changed_files = "" try: - pr_data = await github.get_pull_request(owner, repo, pr_number) - pr_description = pr_data.get("body", "") or "" + pr_description = await _fetch_pr_body(adapter, repo_ref, identity) diff_result = git._run_git( - "diff", "--name-only", f"origin/main...fork/{workspace.branch_name}", check=False + "diff", + "--name-only", + f"origin/main...{push_remote}/{workspace.branch_name}", + check=False, ) changed_files = diff_result.stdout.strip() except Exception as e: @@ -214,13 +240,15 @@ async def rebase_pr(state: WorkflowState) -> WorkflowState: git.stage_all() git.commit(f"[{ticket_key}] merge: resolve conflicts with main") - git.push_to_fork(force=True) - logger.info(f"{ticket_key}: conflicts resolved and pushed to fork") + if use_fork: + git.push_to_fork(force=True) + else: + git.push(force=True, check_conflicts=False) + logger.info(f"{ticket_key}: conflicts resolved and pushed") - await github.create_issue_comment( - owner, - repo, - pr_number, + await adapter.create_comment( + repo_ref, + identity, f"Merge conflicts resolved and pushed. The PR branch has been updated.\n\n" f"Resolved files: {', '.join(f'`{f}`' for f in conflicted_files)}", ) @@ -255,4 +283,3 @@ async def rebase_pr(state: WorkflowState) -> WorkflowState: ) finally: await jira.close() - await github.close() diff --git a/src/forge/workflow/nodes/task_takeover_execution.py b/src/forge/workflow/nodes/task_takeover_execution.py index acac400e3..dd3a7db98 100644 --- a/src/forge/workflow/nodes/task_takeover_execution.py +++ b/src/forge/workflow/nodes/task_takeover_execution.py @@ -11,6 +11,7 @@ PushPersistenceError, build_persistence_error_state, push_to_fork_with_retry, + use_fork_remote, ) from forge.workflow.nodes.workspace_setup import prepare_workspace from forge.workflow.task_takeover.state import TaskTakeoverState @@ -44,7 +45,7 @@ async def execute_task_changes(state: TaskTakeoverState) -> TaskTakeoverState: # Resume safely when another worker cannot see the checkpointed local # workspace. The implementation branch is persisted to the fork # below so a newly cloned workspace contains the reviewed changes. - workspace_path, git = prepare_workspace(state) + workspace_path, git = await prepare_workspace(state) state = {**state, "workspace_path": workspace_path} same_workspace_survived = ( @@ -54,7 +55,7 @@ async def execute_task_changes(state: TaskTakeoverState) -> TaskTakeoverState: ) if state.get("implementation_push_pending") and same_workspace_survived: try: - await push_to_fork_with_retry(git) + await push_to_fork_with_retry(git, use_fork=use_fork_remote(state)) except PushPersistenceError as exc: return cast( TaskTakeoverState, @@ -189,7 +190,7 @@ async def execute_task_changes(state: TaskTakeoverState) -> TaskTakeoverState: # Review may be consumed by another worker with a different local # filesystem. Persist the exact commit before checkpointing this node. try: - await push_to_fork_with_retry(git) + await push_to_fork_with_retry(git, use_fork=use_fork_remote(state)) except PushPersistenceError as exc: pending_state = { **execution_state, diff --git a/src/forge/workflow/nodes/task_takeover_review.py b/src/forge/workflow/nodes/task_takeover_review.py index 68651d2e0..06688f26b 100644 --- a/src/forge/workflow/nodes/task_takeover_review.py +++ b/src/forge/workflow/nodes/task_takeover_review.py @@ -16,6 +16,7 @@ from forge.workflow.nodes.workspace_setup import prepare_workspace from forge.workflow.task_takeover.state import TaskTakeoverState as WorkflowState from forge.workflow.utils import merge_review_exhaustion, update_state_timestamp +from forge.workflow.utils.source_control import get_adapter from forge.workspace.git_ops import GitOperations from forge.workspace.manager import Workspace @@ -69,7 +70,7 @@ async def run_qualitative_review(state: WorkflowState) -> WorkflowState: # A workflow can resume on a different worker from the one that ran # implementation. Never trust the checkpointed local path: recover # the branch from the fork when that path is not visible here. - workspace_path, _ = prepare_workspace(state) + workspace_path, _ = await prepare_workspace(state) state = {**state, "workspace_path": workspace_path} # Fetch ticket details from Jira @@ -78,13 +79,15 @@ async def run_qualitative_review(state: WorkflowState) -> WorkflowState: acceptance_criteria = _extract_acceptance_criteria(description) # Initialize GitOperations to retrieve git diff + repo_ref, adapter = get_adapter(current_repo) git = GitOperations( Workspace( path=Path(workspace_path), repo_name=current_repo, branch_name=state.get("context", {}).get("branch_name", ""), ticket_key=ticket_key, - ) + ), + await adapter.get_git_credentials(repo_ref), ) git_diff = collect_git_diff(git) diff --git a/src/forge/workflow/nodes/workspace_setup.py b/src/forge/workflow/nodes/workspace_setup.py index 8b4470343..e9f66f868 100644 --- a/src/forge/workflow/nodes/workspace_setup.py +++ b/src/forge/workflow/nodes/workspace_setup.py @@ -10,8 +10,8 @@ from typing import Any from forge.config import get_settings -from forge.integrations.github.client import GitHubClient from forge.integrations.jira.client import JiraClient +from forge.integrations.source_control.errors import NotFoundError, ProviderConfigError from forge.workflow.nodes.git_persistence import push_to_fork_with_retry from forge.workflow.utils import update_state_timestamp from forge.workflow.utils.jira_status import ( @@ -19,6 +19,7 @@ set_implementing_label, transition_tasks_to_in_progress, ) +from forge.workflow.utils.source_control import get_adapter from forge.workspace.git_ops import GitOperations from forge.workspace.guardrails import GuardrailsLoader from forge.workspace.handoff import materialize_handoff @@ -56,7 +57,7 @@ def _remove_workspace_backup(path: Path) -> None: time.sleep(_BACKUP_CLEANUP_RETRY_DELAY_SECONDS) -def _recreate_workspace_from_fork( +async def _recreate_workspace_from_fork( *, ticket_key: str, current_repo: str, @@ -66,11 +67,14 @@ def _recreate_workspace_from_fork( state: WorkflowState, stale_workspace_path: str | None = None, ) -> tuple[str, GitOperations]: - if not branch_name or not current_repo or not fork_owner or not fork_repo: + if not branch_name or not current_repo: raise ValueError( f"Cannot recreate workspace for {ticket_key}: " - "missing branch_name, current_repo, fork_owner, or fork_repo in state" + "missing branch_name or current_repo in state" ) + use_fork = bool(fork_owner and fork_repo) + repo_ref, adapter = get_adapter(current_repo) + credentials = await adapter.get_git_credentials(repo_ref) manager = WorkspaceManager(base_dir=get_settings().workspace_base_dir) workspace_obj = manager.create_workspace(repo_name=current_repo, ticket_key=ticket_key) @@ -90,11 +94,14 @@ def _recreate_workspace_from_fork( branch_name=branch_name, ticket_key=ticket_key, ) - git = GitOperations(replacement_workspace) + git = GitOperations(replacement_workspace, credentials) try: git.clone() - git.add_fork_remote(fork_owner, fork_repo) - git.checkout_branch(branch_name, remote="fork") + if use_fork: + git.add_fork_remote(fork_owner, fork_repo) + git.checkout_branch(branch_name, remote="fork") + else: + git.checkout_branch(branch_name, remote="origin") except Exception: shutil.rmtree(replacement_path, ignore_errors=True) raise @@ -142,16 +149,16 @@ def _recreate_workspace_from_fork( return str(target_path), git -def prepare_workspace( +async def prepare_workspace( state: WorkflowState, - remote: str = "fork", ) -> tuple[str, GitOperations]: """Return a workspace path and GitOperations aligned with the remote. If the workspace recorded in state already exists on disk, the branch is rebased onto the remote so that subsequent pushes cannot be rejected as - non-fast-forward. If the workspace is missing it is recreated from the - fork branch via a fresh clone. + non-fast-forward. If the workspace is missing it is recreated via a fresh + clone, from the fork branch when the repo uses fork mode, or from origin + directly when it uses direct mode (state has no fork_owner/fork_repo). This is the single canonical entry point for all implementation nodes (implement_review, attempt_ci_fix, etc.) instead of duplicating @@ -159,7 +166,6 @@ def prepare_workspace( Args: state: Current workflow state. - remote: Remote name to sync with when the workspace exists (default: 'fork'). Returns: Tuple of (workspace_path, GitOperations). @@ -173,7 +179,10 @@ def prepare_workspace( branch_name = state.get("context", {}).get("branch_name", "") fork_owner = state.get("fork_owner", "") fork_repo = state.get("fork_repo", "") + remote = "fork" if (fork_owner and fork_repo) else "origin" ticket_key = state["ticket_key"] + repo_ref, adapter = get_adapter(current_repo) + credentials = await adapter.get_git_credentials(repo_ref) if workspace_path and Path(workspace_path).exists(): workspace = Workspace( @@ -182,7 +191,7 @@ def prepare_workspace( branch_name=branch_name, ticket_key=ticket_key, ) - git = GitOperations(workspace) + git = GitOperations(workspace, credentials) try: git.pull_rebase(remote=remote) except Exception as e: @@ -191,7 +200,7 @@ def prepare_workspace( ticket_key, e, ) - return _recreate_workspace_from_fork( + return await _recreate_workspace_from_fork( ticket_key=ticket_key, current_repo=current_repo, branch_name=branch_name, @@ -203,7 +212,7 @@ def prepare_workspace( return workspace_path, git # Workspace is missing — recreate from fork branch. - return _recreate_workspace_from_fork( + return await _recreate_workspace_from_fork( ticket_key=ticket_key, current_repo=current_repo, branch_name=branch_name, @@ -259,7 +268,51 @@ async def setup_workspace(state: WorkflowState) -> WorkflowState: current_repo = repos[0] # Validate repository name - if current_repo == "unknown" or "/" not in current_repo: + if current_repo == "unknown": + logger.error( + f"Invalid repository name '{current_repo}' for {ticket_key}. " + "Repository must be in 'owner/repo' format." + ) + return { + **state, + "last_error": f"Invalid repository '{current_repo}'. Tasks must specify a valid 'owner/repo' format.", + "current_node": "setup_workspace", + } + + # Resolve current_repo (which may be a repos.yaml alias, not "owner/repo") + # to its canonical namespace before cloning. Downstream PR-state keys and + # webhook lookups are both keyed off this same value, so cloning under the + # raw alias would silently desync CI/review/merge event matching later. + raw_current_repo = current_repo + try: + repo_ref, adapter = get_adapter(current_repo) + except (NotFoundError, ProviderConfigError) as exc: + logger.error(f"Cannot resolve repository '{current_repo}' for {ticket_key}: {exc}") + return { + **state, + "last_error": f"Invalid repository '{current_repo}': {exc}", + "current_node": "setup_workspace", + } + current_repo = repo_ref.namespace + + # tasks_by_repo and repos_to_process were built (by task_router/plan_bug_fix/ + # task_takeover_planning) using the same raw identifier as current_repo. Once + # that identifier is canonicalized above, both must be rewritten in lockstep + # or implementation's tasks_by_repo lookup comes back empty and + # route_after_pr's repos_to_process/repos_completed matching never + # converges (the alias entry can never be marked completed). + repos_to_process = state.get("repos_to_process") + if raw_current_repo != current_repo: + if raw_current_repo in tasks_by_repo: + tasks_by_repo = dict(tasks_by_repo) + alias_tasks = tasks_by_repo.pop(raw_current_repo) + tasks_by_repo[current_repo] = tasks_by_repo.get(current_repo, []) + alias_tasks + if repos_to_process: + repos_to_process = [ + current_repo if r == raw_current_repo else r for r in repos_to_process + ] + + if "/" not in current_repo: logger.error( f"Invalid repository name '{current_repo}' for {ticket_key}. " "Repository must be in 'owner/repo' format." @@ -322,7 +375,7 @@ async def setup_workspace(state: WorkflowState) -> WorkflowState: # Initialize git operations logger.info(f"Initializing git operations for {workspace}") - git = GitOperations(workspace) + git = GitOperations(workspace, await adapter.get_git_credentials(repo_ref)) # Clone repository (600s timeout) logger.info( @@ -336,47 +389,36 @@ async def setup_workspace(state: WorkflowState) -> WorkflowState: # pushes to this fork, so continuing without it would guarantee a # later persistence failure. default_branch = "main" - fork_owner = state.get("fork_owner", "") - fork_repo_name = state.get("fork_repo", "") - if current_repo and "/" in current_repo: - owner, repo_name = current_repo.split("/", 1) - github = GitHubClient() - try: - try: - repo_data = await github.get_repository(owner, repo_name) - default_branch = repo_data.get("default_branch", "main") - logger.info(f"Detected default branch for {current_repo}: {default_branch}") - except Exception as exc: - logger.warning(f"Could not detect default branch for {current_repo}: {exc}") - - if not fork_owner or not fork_repo_name: - fork_data = await github.get_or_create_fork(owner, repo_name) - fork_owner = fork_data["owner"]["login"] - fork_repo_name = fork_data["name"] - - await github.sync_fork_with_upstream( - fork_owner, - fork_repo_name, - branch=default_branch, - ) - finally: - await github.close() - - # Set up feature branch. - git.add_fork_remote(fork_owner, fork_repo_name) - branch_exists_on_fork = git.remote_branch_exists(workspace.branch_name, remote="fork") - if branch_exists_on_fork: + try: + default_branch = await adapter.resolve_default_branch(repo_ref) + logger.info(f"Detected default branch for {current_repo}: {default_branch}") + except Exception as exc: + logger.warning(f"Could not detect default branch for {current_repo}: {exc}") + + write_target = await adapter.ensure_write_target(repo_ref) + fork_owner = write_target.fork_owner or "" + fork_repo_name = write_target.fork_repo or "" + + # Set up feature branch. Direct-mode repos (write_target.push_remote_name + # == "origin") have no fork identity — work directly off origin instead + # of building an invalid fork remote URL from empty owner/repo. + use_fork = bool(fork_owner and fork_repo_name) + push_remote = "fork" if use_fork else "origin" + if use_fork: + git.add_fork_remote(fork_owner, fork_repo_name) + branch_exists_remotely = git.remote_branch_exists(workspace.branch_name, remote=push_remote) + if branch_exists_remotely: logger.info( - f"Branch '{workspace.branch_name}' exists on fork " - f"{fork_owner}/{fork_repo_name} — checking it out" + f"Branch '{workspace.branch_name}' exists on {push_remote} — checking it out" ) - git.checkout_branch(workspace.branch_name, remote="fork") + git.checkout_branch(workspace.branch_name, remote=push_remote) else: git.create_branch(default_branch) # The next graph node may run on a worker that cannot see this # local workspace. Publish the new branch before checkpointing so - # that worker can recreate it from the fork on the first attempt. - await push_to_fork_with_retry(git) + # that worker can recreate it from the fork (or origin, in direct + # mode) on the first attempt. + await push_to_fork_with_retry(git, use_fork=use_fork) # Create .forge directory for task handoff forge_dir = workspace.path / ".forge" @@ -415,18 +457,20 @@ async def setup_workspace(state: WorkflowState) -> WorkflowState: logger.info(f"Workspace ready: {workspace}") - return update_state_timestamp( - { - **state, - "workspace_path": str(workspace.path), - "current_repo": current_repo, - "fork_owner": fork_owner, - "fork_repo": fork_repo_name, - "context": context, - "current_node": "implementation", - "last_error": None, - } - ) + updates: dict[str, Any] = { + **state, + "workspace_path": str(workspace.path), + "current_repo": current_repo, + "tasks_by_repo": tasks_by_repo, + "fork_owner": fork_owner, + "fork_repo": fork_repo_name, + "context": context, + "current_node": "implementation", + "last_error": None, + } + if repos_to_process is not None: + updates["repos_to_process"] = repos_to_process + return update_state_timestamp(updates) except Exception as e: logger.error(f"Workspace setup failed for {ticket_key}: {e}") diff --git a/src/forge/workflow/pr_state.py b/src/forge/workflow/pr_state.py index 3adc734c9..620bb05fc 100644 --- a/src/forge/workflow/pr_state.py +++ b/src/forge/workflow/pr_state.py @@ -3,6 +3,8 @@ from copy import deepcopy from typing import Any, TypedDict +from forge.integrations.source_control.contracts import NormalizedEvent + class PullRequestState(TypedDict, total=False): url: str @@ -20,6 +22,7 @@ class PullRequestState(TypedDict, total=False): review_response_posted: bool merged: bool lifecycle_node: str + pending_ci_event: bool _ACTIVE_FIELDS = { @@ -35,6 +38,7 @@ class PullRequestState(TypedDict, total=False): "review_comments": "review_comments", "contested_comments": "contested_comments", "review_response_posted": "review_response_posted", + "pending_ci_event": "pending_ci_event", } _PR_LIFECYCLE_NODES = { @@ -47,47 +51,83 @@ class PullRequestState(TypedDict, total=False): } -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) +def _numbered_key(repo: str, number: int | str) -> str: + """Per-PR dict key from a repo namespace and a known PR number.""" + return f"{repo}:{number}" + + +def _url_key(repo: str, url: str) -> str: + """Fallback dict key for a PR whose number is not yet known (see module docstring).""" + return f"{repo}:{url}" + + +def _lookup_record( + pull_requests: dict[str, Any], repo: str, number: int | str | None, url: str | None +) -> tuple[str | None, dict[str, Any] | None]: + """Find a PR record for ``repo``, preferring the per-PR numbered key and + falling back to a URL-keyed record for a PR saved before its number was known. + + Also falls back to the legacy bare-``repo`` key from before per-PR keying + was introduced, so a workflow checkpointed mid-CI/mid-review at deploy + time doesn't get stranded — its record still lives under ``repo`` alone + until the next ``save_active_pull_request`` migrates it to a per-PR key. + + Returns ``(key, record)`` or ``(None, None)`` when no record matches. + """ + if number is not None: + record = pull_requests.get(_numbered_key(repo, number)) + if isinstance(record, dict): + return _numbered_key(repo, number), record + if url: + record = pull_requests.get(_url_key(repo, url)) + if isinstance(record, dict): + return _url_key(repo, url), record + record = pull_requests.get(repo) + if isinstance(record, dict): + return repo, record + return None, None + + +def find_active_pull_request( + pull_requests: dict[str, Any], repo: str, number: int | str | None, url: str | None +) -> tuple[str | None, dict[str, Any] | None]: + """Public wrapper around ``_lookup_record`` for callers outside this module + (e.g. ``ci_evaluator``) that need the same numbered-key/URL-fallback lookup + rather than duplicating the key-construction logic.""" + return _lookup_record(pull_requests, repo, number, url) + + +def _record_matches_event(record: dict[str, Any], event: NormalizedEvent) -> bool: + if event.change_request is None: + return False + number = event.change_request.identity.native_id 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") + return record.get("number") is None and event.change_request.url == 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.""" + """Copy the scalar compatibility view into its per-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")): + number = state.get("current_pr_number") + url = state.get("current_pr_url") + if not repo or (number is None and not url): return state + key = _numbered_key(repo, number) if number is not None else _url_key(repo, url) existing_pull_requests = state.get("pull_requests", {}) pull_requests = dict(existing_pull_requests) - record = deepcopy(existing_pull_requests.get(repo, {})) + # Look up by number-or-url rather than `key` alone: a record saved before + # the PR number was known lives under the url key, and once the number + # becomes available `key` switches to the numbered key. Without this + # lookup that stale url-keyed record would be missed and a fresh, empty + # record created in its place, orphaning fields like lifecycle_node. + existing_key, existing_record = _lookup_record(pull_requests, repo, number, url) + record = deepcopy(existing_record) if existing_record is not None else {} + if existing_key is not None and existing_key != key: + del pull_requests[existing_key] for scalar, per_pr in _ACTIVE_FIELDS.items(): if scalar in state: record[per_pr] = state[scalar] @@ -99,27 +139,34 @@ def save_active_pull_request(state: dict[str, Any]) -> dict[str, Any]: # Default to the post-PR entry node (teardown_workspace → human_review_gate) # so a stale record resolves to a real node on resume rather than a removed one. record.setdefault("lifecycle_node", "human_review_gate") - pull_requests[repo] = record + pull_requests[key] = record return {**state, "pull_requests": pull_requests} def activate_pull_request_for_event( - state: dict[str, Any], payload: dict[str, Any] + state: dict[str, Any], event: NormalizedEvent | None ) -> 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) + """Select the PR targeted by a source-control webhook as the scalar + compatibility view. ``event`` is ``None`` for a Jira message (no PR to + activate), in which case ``state`` is returned unchanged.""" + if event is None or event.change_request is None: + return state + + repo = event.repo_ref.namespace + number = event.change_request.identity.native_id + url = event.change_request.url 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): + key, record = _lookup_record(pull_requests, repo, number, url) + if record is None or not _record_matches_event(record, event): 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 + record = updated_pull_requests.pop(key) + record["number"] = number + updated_pull_requests[_numbered_key(repo, number)] = record 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(): @@ -142,20 +189,31 @@ def all_pull_requests_merged(state: dict[str, Any]) -> bool: def mark_active_pull_request_merged(state: dict[str, Any]) -> dict[str, Any]: - """Mark the selected per-repository PR merged without changing other records.""" + """Mark the selected per-PR record merged without changing other records.""" repo = state.get("current_repo") + if not repo: + return state 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): + key, record = _lookup_record( + existing_pull_requests, repo, state.get("current_pr_number"), state.get("current_pr_url") + ) + if record is None: return state + pull_requests = dict(existing_pull_requests) + record = deepcopy(record) record["merged"] = True - pull_requests[repo] = record + pull_requests[key] = record return {**state, "pull_requests": pull_requests} -def event_targets_pull_request(state: dict[str, Any], payload: dict[str, Any]) -> bool: +def event_targets_pull_request(state: dict[str, Any], event: NormalizedEvent | None) -> 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) + if event is None or event.change_request is None: + return False + _, record = _lookup_record( + state.get("pull_requests", {}), + event.repo_ref.namespace, + event.change_request.identity.native_id, + event.change_request.url, + ) + return record is not None and _record_matches_event(record, event) diff --git a/src/forge/workflow/utils/__init__.py b/src/forge/workflow/utils/__init__.py index 2258bf7d7..426e09093 100644 --- a/src/forge/workflow/utils/__init__.py +++ b/src/forge/workflow/utils/__init__.py @@ -34,6 +34,11 @@ "ci_evaluator": "ci_evaluator", "attempt_ci_fix": "ci_evaluator", "rebase_pr": "rebase_pr", + # wait_for_ci_gate was merged into human_review_gate so CI and review run + # concurrently; this compatibility alias lets a ticket already + # checkpointed at wait_for_ci_gate before that merge resume correctly + # instead of silently restarting its whole workflow from the beginning. + "wait_for_ci_gate": "human_review_gate", } _TERMINAL_NODES: frozenset[str] = frozenset({"complete"}) diff --git a/src/forge/workflow/utils/proposal_review_threads.py b/src/forge/workflow/utils/proposal_review_threads.py index 4d70dd549..80cbf1214 100644 --- a/src/forge/workflow/utils/proposal_review_threads.py +++ b/src/forge/workflow/utils/proposal_review_threads.py @@ -116,7 +116,7 @@ async def reply_to_proposal_decisions( if not str(decision.get("response", "")).strip(): decision["response"] = default_response await reply_to_review_decisions( - repo_full_name=repo_full_name, + current_repo=repo_full_name, pr_number=pr_number, decisions=decisions, dispositions=dispositions, diff --git a/src/forge/workflow/utils/review_decisions.py b/src/forge/workflow/utils/review_decisions.py index f71f5dc46..53b51470e 100644 --- a/src/forge/workflow/utils/review_decisions.py +++ b/src/forge/workflow/utils/review_decisions.py @@ -3,6 +3,8 @@ import logging from typing import Any +from forge.workflow.utils.source_control import get_adapter, identity_for + logger = logging.getLogger(__name__) @@ -38,45 +40,47 @@ def flatten_review_threads(threads: list[dict[str, Any]]) -> list[dict[str, Any] ] -def decision_matches_comment(decision: dict[str, Any], comment_id: int) -> bool: +def decision_matches_comment(decision: dict[str, Any], comment_id: str | int) -> bool: """Match either the reviewer comment or Forge's reply in the same thread.""" - return comment_id in (decision.get("comment_id"), decision.get("forge_reply_id")) + target = str(comment_id) + candidates = { + str(value) + for value in (decision.get("comment_id"), decision.get("forge_reply_id")) + if value is not None + } + return target in candidates async def reply_to_review_decisions( *, - repo_full_name: str, + current_repo: str, pr_number: int | None, decisions: list[dict[str, Any]], dispositions: set[str] | None = None, skip_addressed: bool = False, ) -> None: """Reply consistently and retain Forge's reply ID for later correlation.""" - if not repo_full_name or "/" not in repo_full_name or not pr_number or not decisions: + if not current_repo or not pr_number or not decisions: return - from forge.integrations.github.client import GitHubClient - - owner, repo = repo_full_name.split("/", 1) - github = GitHubClient() try: - for decision in decisions: - if dispositions is not None and decision.get("disposition") not in dispositions: - continue - if skip_addressed and decision.get("status") == "addressed": - continue - comment_id = decision.get("comment_id") - response = str(decision.get("response", "")).strip() - if not isinstance(comment_id, int) or not response: - continue - try: - reply = await github.reply_to_review_comment( - owner, repo, pr_number, comment_id, response - ) - reply_id = reply.get("id") if isinstance(reply, dict) else None - if isinstance(reply_id, int): - decision["forge_reply_id"] = reply_id - except Exception as exc: - logger.warning("Failed replying to review comment %s: %s", comment_id, exc) - finally: - await github.close() + repo_ref, adapter = get_adapter(current_repo) + identity = identity_for(repo_ref, pr_number) + except Exception as exc: + logger.warning("Could not resolve adapter for %s: %s", current_repo, exc) + return + + for decision in decisions: + if dispositions is not None and decision.get("disposition") not in dispositions: + continue + if skip_addressed and decision.get("status") == "addressed": + continue + comment_id = decision.get("comment_id") + response = str(decision.get("response", "")).strip() + if comment_id is None or not response: + continue + try: + reply = await adapter.reply_to_comment(repo_ref, identity, str(comment_id), response) + decision["forge_reply_id"] = reply.id + except Exception as exc: + logger.warning("Failed replying to review comment %s: %s", comment_id, exc) diff --git a/src/forge/workflow/utils/source_control.py b/src/forge/workflow/utils/source_control.py new file mode 100644 index 000000000..ffdcf34e8 --- /dev/null +++ b/src/forge/workflow/utils/source_control.py @@ -0,0 +1,49 @@ +"""Node-facing bridge from a repo identifier to a provider adapter. + +Wraps the registry so workflow nodes never import a concrete client or +re-derive the resolve() pattern. Local-git plumbing (GitOperations) is +unchanged; only API calls go through the returned adapter. +""" + +from forge.integrations.source_control.contracts import ( + ChangeRequestIdentity, + Provider, + RepositoryRef, + ResolvedRepository, + SourceControlProvider, +) +from forge.integrations.source_control.errors import NotFoundError +from forge.integrations.source_control.registry import get_registry + + +def resolve_repository( + identifier: str, provider_hint: Provider | None = None +) -> ResolvedRepository: + """Resolve a repo identifier (a `repos.yaml` id or a bare `owner/repo`).""" + if not identifier: + raise NotFoundError("Cannot resolve an empty repository identifier") + return get_registry().resolve(identifier, provider_hint=provider_hint) + + +def get_adapter(identifier: str) -> tuple[RepositoryRef, SourceControlProvider]: + """Resolve `identifier` to its (RepositoryRef, adapter) pair. + + Raises NotFoundError when the identifier is empty or the provider has no + registered adapter factory (which would leave `adapter` None). + """ + resolved = resolve_repository(identifier) + if resolved.adapter is None: + raise NotFoundError( + f"'{identifier}' resolved to provider '{resolved.repo_ref.provider}' " + "with no registered adapter" + ) + return resolved.repo_ref, resolved.adapter + + +def identity_for(repo_ref: RepositoryRef, native_id: str | int | None) -> ChangeRequestIdentity: + """Build the composite change-request identity for a repo + native PR/MR id.""" + return ChangeRequestIdentity( + connection=repo_ref.connection, + repository_id=repo_ref.id, + native_id=native_id, + ) diff --git a/src/forge/workspace/git_ops.py b/src/forge/workspace/git_ops.py index feafaebd9..bd1a26868 100644 --- a/src/forge/workspace/git_ops.py +++ b/src/forge/workspace/git_ops.py @@ -1,10 +1,12 @@ """Git operations for workspace management.""" import logging +import os import subprocess from pathlib import Path from forge.config import get_settings +from forge.integrations.source_control.contracts import GitCredentials from forge.utils.redaction import redact_secrets from forge.workspace.manager import Workspace @@ -14,13 +16,20 @@ class GitOperations: """Git operations for cloning, branching, committing, and pushing.""" - def __init__(self, workspace: Workspace): + def __init__(self, workspace: Workspace, credentials: GitCredentials): """Initialize git operations for a workspace. Args: workspace: Workspace to operate on. + credentials: Host/token/CA for this workspace's connection. Every + clone/remote URL this class builds is derived from these, + never from process-wide settings, so operations against a + non-default connection (e.g. GitHub Enterprise, or a second + org with its own token) hit the right host with the right + credential. """ self.workspace = workspace + self.credentials = credentials self.settings = get_settings() # Set by workspace recovery when this instance represents a replacement # clone rather than the workspace recorded in workflow state. The path @@ -33,6 +42,22 @@ def repo_path(self) -> Path: """Get the repository path.""" return self.workspace.path + def _remote_url(self, owner: str, repo: str) -> str: + """Build an authenticated HTTPS clone/remote URL for owner/repo on + this workspace's connection host.""" + return f"https://x-access-token:{self.credentials.token}@{self.credentials.host}/{owner}/{repo}.git" + + def _git_env(self) -> dict[str, str] | None: + """Subprocess environment for git commands, trusting this connection's + CA bundle when it has one (self-signed GitHub Enterprise Server certs). + + Returns None (inherit the process environment unmodified) when no + ca_path is configured -- the common case. + """ + if not self.credentials.ca_path: + return None + return {**os.environ, "GIT_SSL_CAINFO": self.credentials.ca_path} + def _run_git( self, *args: str, @@ -59,6 +84,7 @@ def _run_git( capture_output=capture_output, text=True, check=False, + env=self._git_env(), ) if check and result.returncode != 0: @@ -75,12 +101,13 @@ def clone( """Clone the repository into the workspace. Args: - repo_url: Repository URL. Constructs from settings if None. + repo_url: Repository URL. Built from this workspace's connection + credentials if None. timeout: Timeout in seconds for the clone operation (default 600s). """ if repo_url is None: - token = self.settings.github_token.get_secret_value() - repo_url = f"https://x-access-token:{token}@github.com/{self.workspace.repo_name}.git" + owner, _, repo = self.workspace.repo_name.partition("/") + repo_url = self._remote_url(owner, repo) # Build clone command (single-branch for faster clone) cmd = ["git", "clone", "--single-branch", repo_url, str(self.repo_path)] @@ -99,6 +126,7 @@ def clone( text=True, check=True, timeout=timeout, + env=self._git_env(), ) elapsed = time.time() - start_time logger.info(f"Clone completed for {self.workspace.repo_name} in {elapsed:.1f}s") @@ -135,7 +163,7 @@ def pull_rebase(self, remote: str = "fork") -> None: """ branch = self.workspace.branch_name logger.info(f"Syncing with {remote}/{branch} before implementing changes") - self._run_git("fetch", remote) + self._run_git("fetch", remote, f"{branch}:refs/remotes/{remote}/{branch}", check=False) if not self.remote_branch_exists(branch, remote=remote): logger.info( "Remote branch %s/%s does not exist yet; skipping rebase before first push", @@ -153,8 +181,7 @@ def add_fork_remote(self, fork_owner: str, fork_repo: str) -> None: fork_owner: Owner of the fork repository. fork_repo: Name of the fork repository. """ - token = self.settings.github_token.get_secret_value() - fork_url = f"https://x-access-token:{token}@github.com/{fork_owner}/{fork_repo}.git" + fork_url = self._remote_url(fork_owner, fork_repo) # Check if remote already exists result = self._run_git("remote", check=False) @@ -214,6 +241,7 @@ def checkout_branch(self, branch_name: str | None = None, remote: str = "origin" # create it tracking the specified remote. result = self._run_git("checkout", branch, check=False) if result.returncode != 0: + self._run_git("fetch", remote, f"{branch}:refs/remotes/{remote}/{branch}", check=False) self._run_git("checkout", "-b", branch, f"{remote}/{branch}") logger.info(f"Checked out branch {branch}") diff --git a/tests/conftest.py b/tests/conftest.py index d5a313d1e..1cab8332e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -7,6 +7,9 @@ os.environ.setdefault("JIRA_API_TOKEN", "test-token") os.environ.setdefault("JIRA_USER_EMAIL", "test@example.com") os.environ.setdefault("GITHUB_TOKEN", "test-github-token") +os.environ.setdefault("LLM_BACKEND", "anthropic") +os.environ.setdefault("LLM_MODEL", "claude-3-5-sonnet-20241022") +os.environ.setdefault("ANTHROPIC_API_KEY", "test-key") from collections.abc import AsyncGenerator, Generator from pathlib import Path diff --git a/tests/contracts/fixtures/github_api_responses/check_run_failure.json b/tests/contracts/fixtures/github_api_responses/check_run_failure.json deleted file mode 100644 index 78996b0ff..000000000 --- a/tests/contracts/fixtures/github_api_responses/check_run_failure.json +++ /dev/null @@ -1,50 +0,0 @@ -{ - "action": "completed", - "check_run": { - "id": 29487615893, - "name": "CI / Tests", - "head_sha": "def456abc789012345678901234567890123", - "status": "completed", - "conclusion": "failure", - "started_at": "2024-03-20T11:00:00Z", - "completed_at": "2024-03-20T11:02:45Z", - "output": { - "title": "Tests failed", - "summary": "3 of 142 tests failed", - "text": "## Failed Tests\n\n### test_login_validation\n```\nAssertionError: Expected status 200, got 400\n File \"tests/test_auth.py\", line 45\n assert response.status_code == 200\n```\n\n### test_password_special_chars\n```\nValueError: Invalid character in password\n File \"src/auth/validators.py\", line 23\n```", - "annotations_count": 3 - }, - "check_suite": { - "id": 21234567891, - "head_branch": "bugfix/PROJ-156-password-chars", - "status": "completed", - "conclusion": "failure" - }, - "pull_requests": [ - { - "url": "https://api.github.com/repos/acme/backend/pulls/45", - "id": 1789012346, - "number": 45, - "head": { - "ref": "bugfix/PROJ-156-password-chars", - "sha": "def456abc789012345678901234567890123" - }, - "base": { - "ref": "main" - } - } - ] - }, - "repository": { - "id": 123456789, - "name": "backend", - "full_name": "acme/backend", - "owner": { - "login": "acme" - } - }, - "sender": { - "login": "github-actions[bot]", - "id": 41898282 - } -} diff --git a/tests/contracts/fixtures/github_api_responses/check_run_success.json b/tests/contracts/fixtures/github_api_responses/check_run_success.json deleted file mode 100644 index 2b32ac4dc..000000000 --- a/tests/contracts/fixtures/github_api_responses/check_run_success.json +++ /dev/null @@ -1,73 +0,0 @@ -{ - "action": "completed", - "check_run": { - "id": 29487615892, - "name": "CI / Tests", - "node_id": "CR_kwDOKz1234567890", - "head_sha": "abc123def456789012345678901234567890", - "external_id": "", - "url": "https://api.github.com/repos/acme/backend/check-runs/29487615892", - "html_url": "https://github.com/acme/backend/runs/29487615892", - "details_url": "https://github.com/acme/backend/actions/runs/12345678901", - "status": "completed", - "conclusion": "success", - "started_at": "2024-03-20T10:00:00Z", - "completed_at": "2024-03-20T10:05:32Z", - "output": { - "title": "Tests passed", - "summary": "All 142 tests passed in 5m 32s", - "text": null, - "annotations_count": 0, - "annotations_url": "https://api.github.com/repos/acme/backend/check-runs/29487615892/annotations" - }, - "check_suite": { - "id": 21234567890, - "head_branch": "feature/PROJ-104-oauth", - "head_sha": "abc123def456789012345678901234567890", - "status": "completed", - "conclusion": "success" - }, - "app": { - "id": 15368, - "slug": "github-actions", - "name": "GitHub Actions" - }, - "pull_requests": [ - { - "url": "https://api.github.com/repos/acme/backend/pulls/42", - "id": 1789012345, - "number": 42, - "head": { - "ref": "feature/PROJ-104-oauth", - "sha": "abc123def456789012345678901234567890", - "repo": { - "id": 123456789, - "name": "backend", - "full_name": "acme/backend" - } - }, - "base": { - "ref": "main", - "sha": "main123456789012345678901234567890" - } - } - ] - }, - "repository": { - "id": 123456789, - "name": "backend", - "full_name": "acme/backend", - "private": true, - "owner": { - "login": "acme", - "id": 12345678, - "type": "Organization" - }, - "default_branch": "main" - }, - "sender": { - "login": "github-actions[bot]", - "id": 41898282, - "type": "Bot" - } -} diff --git a/tests/contracts/fixtures/github_api_responses/pull_request_opened.json b/tests/contracts/fixtures/github_api_responses/pull_request_opened.json deleted file mode 100644 index 31605aa31..000000000 --- a/tests/contracts/fixtures/github_api_responses/pull_request_opened.json +++ /dev/null @@ -1,70 +0,0 @@ -{ - "action": "opened", - "number": 42, - "pull_request": { - "id": 1789012345, - "number": 42, - "state": "open", - "locked": false, - "title": "PROJ-104: Implement OAuth2 authentication flow", - "body": "## Summary\n\n- Added Google OAuth2 provider\n- Implemented token refresh logic\n- Added secure token storage\n\n## Test Plan\n\n- [x] Unit tests added\n- [x] Integration tests pass locally\n\nCloses PROJ-104", - "created_at": "2024-03-20T10:00:00Z", - "updated_at": "2024-03-20T10:00:00Z", - "closed_at": null, - "merged_at": null, - "merge_commit_sha": null, - "head": { - "label": "acme:feature/PROJ-104-oauth", - "ref": "feature/PROJ-104-oauth", - "sha": "abc123def456789012345678901234567890", - "user": { - "login": "forge-bot", - "id": 98765432 - }, - "repo": { - "id": 123456789, - "name": "backend", - "full_name": "acme/backend" - } - }, - "base": { - "label": "acme:main", - "ref": "main", - "sha": "main123456789012345678901234567890", - "repo": { - "id": 123456789, - "name": "backend", - "full_name": "acme/backend" - } - }, - "html_url": "https://github.com/acme/backend/pull/42", - "diff_url": "https://github.com/acme/backend/pull/42.diff", - "patch_url": "https://github.com/acme/backend/pull/42.patch", - "user": { - "login": "forge-bot", - "id": 98765432, - "type": "Bot" - }, - "draft": false, - "mergeable": null, - "mergeable_state": "unknown", - "merged": false, - "commits": 3, - "additions": 452, - "deletions": 12, - "changed_files": 8 - }, - "repository": { - "id": 123456789, - "name": "backend", - "full_name": "acme/backend", - "owner": { - "login": "acme" - }, - "default_branch": "main" - }, - "sender": { - "login": "forge-bot", - "id": 98765432 - } -} diff --git a/tests/contracts/fixtures/github_api_responses/pull_request_review_approved.json b/tests/contracts/fixtures/github_api_responses/pull_request_review_approved.json deleted file mode 100644 index 5ee3b8f0c..000000000 --- a/tests/contracts/fixtures/github_api_responses/pull_request_review_approved.json +++ /dev/null @@ -1,48 +0,0 @@ -{ - "action": "submitted", - "review": { - "id": 1876543210, - "node_id": "PRR_kwDOKz123456789", - "user": { - "login": "senior-dev", - "id": 11223344, - "type": "User" - }, - "body": "LGTM! Great implementation of the OAuth flow. Just a few minor suggestions but nothing blocking.", - "commit_id": "abc123def456789012345678901234567890", - "submitted_at": "2024-03-20T14:30:00Z", - "state": "approved", - "html_url": "https://github.com/acme/backend/pull/42#pullrequestreview-1876543210", - "pull_request_url": "https://api.github.com/repos/acme/backend/pulls/42", - "author_association": "MEMBER" - }, - "pull_request": { - "id": 1789012345, - "number": 42, - "state": "open", - "title": "PROJ-104: Implement OAuth2 authentication flow", - "head": { - "ref": "feature/PROJ-104-oauth", - "sha": "abc123def456789012345678901234567890" - }, - "base": { - "ref": "main" - }, - "html_url": "https://github.com/acme/backend/pull/42", - "user": { - "login": "forge-bot" - } - }, - "repository": { - "id": 123456789, - "name": "backend", - "full_name": "acme/backend", - "owner": { - "login": "acme" - } - }, - "sender": { - "login": "senior-dev", - "id": 11223344 - } -} diff --git a/tests/contracts/source_control/__init__.py b/tests/contracts/source_control/__init__.py new file mode 100644 index 000000000..d329eb160 --- /dev/null +++ b/tests/contracts/source_control/__init__.py @@ -0,0 +1 @@ +"""Source control provider contract conformance tests.""" diff --git a/tests/contracts/source_control/conformance_suite.py b/tests/contracts/source_control/conformance_suite.py new file mode 100644 index 000000000..8c88a8374 --- /dev/null +++ b/tests/contracts/source_control/conformance_suite.py @@ -0,0 +1,80 @@ +"""Provider-agnostic conformance tests for SourceControlProvider. + +Each function asserts a subset of the protocol contract. Test modules +parameterize these over concrete provider fixtures. +""" +from typing import Any + +from forge.integrations.source_control.contracts import ( + EventKind, + RepositoryRef, + SourceControlProvider, +) + + +async def assert_webhook_verification( + adapter: SourceControlProvider, + valid_headers: dict[str, str], + valid_body: bytes, + invalid_signature_headers: dict[str, str], +) -> None: + """Assert webhook signature verification works correctly. + + Args: + adapter: Provider adapter under test + valid_headers: Headers with correct signature + valid_body: Request body that matches the signature + invalid_signature_headers: Headers with wrong signature + """ + # Valid signature should return True + assert await adapter.verify_webhook(valid_headers, valid_body) + + # Invalid signature should return False + assert not await adapter.verify_webhook(invalid_signature_headers, valid_body) + + +async def assert_webhook_parsing( + adapter: SourceControlProvider, + headers: dict[str, str], + body: bytes, + resolver: Any, + expected_kind: EventKind, + expected_repo_namespace: str, +) -> None: + """Assert webhook parsing produces correct NormalizedEvent. + + Args: + adapter: Provider adapter under test + headers: Webhook request headers + body: Webhook request body + resolver: Registry resolver (or mock) + expected_kind: Expected EventKind value + expected_repo_namespace: Expected repository namespace + """ + event = await adapter.parse_webhook(headers, body, resolver) + + assert event.kind == expected_kind + assert event.repo_ref.namespace == expected_repo_namespace + assert event.actor.login + assert event.received_at is not None + + +async def assert_repository_operations( + adapter: SourceControlProvider, + repo_ref: RepositoryRef, +) -> None: + """Assert repository metadata operations work. + + Args: + adapter: Provider adapter under test + repo_ref: Repository to test against + """ + # Should resolve default branch + branch = await adapter.resolve_default_branch(repo_ref) + assert isinstance(branch, str) + assert len(branch) > 0 + + # Should get authenticated identity + identity = await adapter.get_authenticated_identity(repo_ref) + assert identity.login + assert isinstance(identity.is_bot, bool) diff --git a/tests/contracts/source_control/conftest.py b/tests/contracts/source_control/conftest.py new file mode 100644 index 000000000..0ec21b49a --- /dev/null +++ b/tests/contracts/source_control/conftest.py @@ -0,0 +1,5 @@ +"""Pytest fixtures for source control conformance tests. + +Fixtures for provider adapters will be defined here as the conformance suite +is expanded across tasks. +""" diff --git a/tests/contracts/source_control/test_github_adapter.py b/tests/contracts/source_control/test_github_adapter.py new file mode 100644 index 000000000..99cc333fd --- /dev/null +++ b/tests/contracts/source_control/test_github_adapter.py @@ -0,0 +1,2080 @@ +"""Tests for GitHub source control adapter.""" + +import base64 +import hashlib +import hmac +import json +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest + +from forge.integrations.github.client import GitHubClient, PullRequestCreationResult +from forge.integrations.source_control.contracts import ( + ChangeRequestIdentity, + ChangeRequestState, + CheckConclusion, + CheckRun, + CheckStatus, + Connection, + EventKind, + Provider, + RepositoryRef, + ResolvedRepository, + ReviewState, + WriteTarget, +) +from forge.integrations.source_control.errors import ( + AuthenticationError, + ConflictError, + NotFoundError, + RateLimitedError, + SourceControlError, + TransientProviderError, +) +from forge.integrations.source_control.github.adapter import GitHubAdapter +from tests.contracts.source_control.conformance_suite import ( + assert_repository_operations, + assert_webhook_parsing, + assert_webhook_verification, +) + + +@pytest.fixture +def github_connection() -> Connection: + return Connection( + name="test-github", + provider=Provider.GITHUB, + base_url="https://api.github.com", + credential_env="GITHUB_TOKEN", + webhook_secret_env="GITHUB_WEBHOOK_SECRET", + ) + + +@pytest.fixture +def github_repo_ref() -> RepositoryRef: + return RepositoryRef( + id="test/repo", + provider=Provider.GITHUB, + connection="test-github", + namespace="test/repo", + default_branch="main", + change_request_mode="fork", + ) + + +@pytest.fixture +def github_adapter(github_connection: Connection) -> GitHubAdapter: + return GitHubAdapter(github_connection, credential="test-token-123") + + +@pytest.fixture +def webhook_secret() -> str: + return "test-webhook-secret" + + +@pytest.fixture +def mock_github_http_client(mock_settings) -> GitHubClient: + """A GitHubClient whose underlying httpx.AsyncClient is mocked out.""" + client = GitHubClient(settings=mock_settings) + client._client = AsyncMock(spec=httpx.AsyncClient) + client._client.is_closed = False + return client + + +@pytest.fixture +def github_adapter_with_mock_client( + github_connection: Connection, mock_github_http_client: GitHubClient +) -> GitHubAdapter: + """GitHubAdapter wired up with a mocked GitHubClient for HTTP-free tests.""" + return GitHubAdapter( + github_connection, + credential="test-token-123", + client=mock_github_http_client, + ) + + +@pytest.fixture +def valid_pr_opened_payload() -> dict: + return { + "action": "opened", + "pull_request": { + "number": 42, + "html_url": "https://github.com/test/repo/pull/42", + "title": "Test PR", + "body": "Test body", + "state": "open", + "draft": False, + "head": {"ref": "feature-branch"}, + "base": {"ref": "main"}, + }, + "repository": { + "full_name": "test/repo", + }, + "sender": { + "login": "testuser", + "type": "User", + }, + } + + +def sign_webhook(payload: dict, secret: str) -> str: + """Create GitHub webhook signature.""" + body = json.dumps(payload).encode() + signature = hmac.new(secret.encode(), body, hashlib.sha256).hexdigest() + return f"sha256={signature}" + + +class MockResolver: + """Mock repository resolver for testing.""" + + def __init__(self, repo_ref: RepositoryRef, connection: Connection): + self._repo_ref = repo_ref + self._connection = connection + + def resolve( + self, + identifier: str, # noqa: ARG002 + provider_hint: Provider | None = None, # noqa: ARG002 + ) -> ResolvedRepository: + return ResolvedRepository( + repo_ref=self._repo_ref, + connection=self._connection, + adapter=None, + ) + + +@pytest.mark.asyncio +async def test_webhook_verification(github_adapter: GitHubAdapter, webhook_secret: str): + """Test webhook signature verification using conformance suite.""" + payload = {"test": "data"} + body = json.dumps(payload).encode() + + valid_headers = { + "X-Hub-Signature-256": sign_webhook(payload, webhook_secret), + } + + invalid_headers = { + "X-Hub-Signature-256": "sha256=invalid", + } + + # Pass webhook_secret to adapter for verification + adapter_with_secret = GitHubAdapter( + github_adapter._connection, + credential="test-token", + webhook_secret=webhook_secret, + ) + + await assert_webhook_verification(adapter_with_secret, valid_headers, body, invalid_headers) + + +@pytest.mark.asyncio +async def test_webhook_verification_is_disabled_when_secret_unset( + github_adapter: GitHubAdapter, +): + """An adapter with no webhook secret accepts unsigned synthetic webhooks.""" + payload = {"test": "data"} + body = json.dumps(payload).encode() + + assert github_adapter._webhook_secret is None + assert await github_adapter.verify_webhook({}, body) is True + + +@pytest.mark.asyncio +async def test_webhook_parsing_pr_opened( + github_adapter: GitHubAdapter, + github_repo_ref: RepositoryRef, + github_connection: Connection, + valid_pr_opened_payload: dict, +): + """Test parsing pull_request opened webhook.""" + body = json.dumps(valid_pr_opened_payload).encode() + headers = { + "X-GitHub-Event": "pull_request", + "X-GitHub-Delivery": "test-delivery-123", + } + + resolver = MockResolver(github_repo_ref, github_connection) + + await assert_webhook_parsing( + github_adapter, + headers, + body, + resolver, + expected_kind=EventKind.CR_OPENED, + expected_repo_namespace="test/repo", + ) + + +class TestGetClientCredentialThreading: + """_get_client() must thread self._credential into the constructed client.""" + + def test_uses_injected_client_over_credential( + self, + github_connection: Connection, + mock_github_http_client: GitHubClient, + ): + adapter = GitHubAdapter( + github_connection, + credential="ignored-because-client-injected", + client=mock_github_http_client, + ) + + assert adapter._get_client() is mock_github_http_client + + def test_builds_client_using_configured_credential(self, github_connection: Connection): + adapter = GitHubAdapter(github_connection, credential="per-connection-token") + + client = adapter._get_client() + + assert isinstance(client, GitHubClient) + assert client.settings.github_token.get_secret_value() == "per-connection-token" + + def test_falls_back_to_process_settings_when_no_credential(self, github_connection: Connection): + from forge.config import get_settings + + adapter = GitHubAdapter(github_connection) + + client = adapter._get_client() + + assert isinstance(client, GitHubClient) + assert ( + client.settings.github_token.get_secret_value() + == get_settings().github_token.get_secret_value() + ) + + def test_lazily_constructed_client_is_cached(self, github_connection: Connection): + adapter = GitHubAdapter(github_connection, credential="token-a") + + first = adapter._get_client() + second = adapter._get_client() + + assert first is second + + +class TestResolveDefaultBranch: + @pytest.mark.asyncio + async def test_returns_default_branch_from_api( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + mock_client = mock_github_http_client._client + response = MagicMock() + response.json.return_value = {"default_branch": "develop", "full_name": "test/repo"} + response.raise_for_status = MagicMock() + mock_client.get = AsyncMock(return_value=response) + + branch = await github_adapter_with_mock_client.resolve_default_branch(github_repo_ref) + + mock_client.get.assert_called_once_with("/repos/test/repo") + assert branch == "develop" + + @pytest.mark.asyncio + async def test_falls_back_to_main_when_field_missing( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + mock_client = mock_github_http_client._client + response = MagicMock() + response.json.return_value = {"full_name": "test/repo"} + response.raise_for_status = MagicMock() + mock_client.get = AsyncMock(return_value=response) + + branch = await github_adapter_with_mock_client.resolve_default_branch(github_repo_ref) + + assert branch == "main" + + +class TestTranslateProviderErrors: + """Direct coverage for the `_translate_provider_errors` boundary decorator + applied to every adapter method that calls the GitHub API.""" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("status_code", "expected_exception"), + [ + (401, AuthenticationError), + (403, AuthenticationError), + (429, RateLimitedError), + (500, TransientProviderError), + (503, TransientProviderError), + ], + ) + async def test_status_code_maps_to_neutral_exception( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + status_code: int, + expected_exception: type[Exception], + ): + response = httpx.Response( + status_code, + headers={"Retry-After": "30"} if status_code == 429 else {}, + request=httpx.Request("GET", "https://api.github.com/repos/test/repo"), + ) + mock_github_http_client.get_repository = AsyncMock( + side_effect=httpx.HTTPStatusError("boom", request=response.request, response=response) + ) + + with pytest.raises(expected_exception): + await github_adapter_with_mock_client.resolve_default_branch(github_repo_ref) + + @pytest.mark.asyncio + async def test_rate_limit_parses_retry_after( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + response = httpx.Response( + 429, + headers={"Retry-After": "30"}, + request=httpx.Request("GET", "https://api.github.com/repos/test/repo"), + ) + mock_github_http_client.get_repository = AsyncMock( + side_effect=httpx.HTTPStatusError("boom", request=response.request, response=response) + ) + + with pytest.raises(RateLimitedError) as exc_info: + await github_adapter_with_mock_client.resolve_default_branch(github_repo_ref) + + assert exc_info.value.retry_after == 30.0 + + @pytest.mark.asyncio + async def test_rate_limit_with_non_numeric_retry_after_falls_back_to_none( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """An HTTP-date Retry-After value (valid per RFC 7231) must not crash + the error-translation path with an unhandled ValueError.""" + response = httpx.Response( + 429, + headers={"Retry-After": "Wed, 21 Oct 2026 07:28:00 GMT"}, + request=httpx.Request("GET", "https://api.github.com/repos/test/repo"), + ) + mock_github_http_client.get_repository = AsyncMock( + side_effect=httpx.HTTPStatusError("boom", request=response.request, response=response) + ) + + with pytest.raises(RateLimitedError) as exc_info: + await github_adapter_with_mock_client.resolve_default_branch(github_repo_ref) + + assert exc_info.value.retry_after is None + + @pytest.mark.asyncio + async def test_network_failure_maps_to_transient_provider_error( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + mock_github_http_client.get_repository = AsyncMock( + side_effect=httpx.ConnectTimeout("connection timed out") + ) + + with pytest.raises(TransientProviderError): + await github_adapter_with_mock_client.resolve_default_branch(github_repo_ref) + + @pytest.mark.asyncio + async def test_not_found_status_propagates_unchanged( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """404 has no generic neutral mapping and must be left for callers to + handle themselves rather than being swallowed or re-mapped.""" + response = httpx.Response( + 404, request=httpx.Request("GET", "https://api.github.com/repos/test/repo") + ) + mock_github_http_client.get_repository = AsyncMock( + side_effect=httpx.HTTPStatusError("boom", request=response.request, response=response) + ) + + with pytest.raises(httpx.HTTPStatusError): + await github_adapter_with_mock_client.resolve_default_branch(github_repo_ref) + + +class TestGetAuthenticatedIdentity: + @pytest.mark.asyncio + async def test_returns_human_actor( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + mock_client = mock_github_http_client._client + response = MagicMock() + response.json.return_value = {"login": "octocat", "type": "User"} + response.raise_for_status = MagicMock() + mock_client.get = AsyncMock(return_value=response) + + actor = await github_adapter_with_mock_client.get_authenticated_identity(github_repo_ref) + + mock_client.get.assert_called_once_with("/user") + assert actor.login == "octocat" + assert actor.is_bot is False + + @pytest.mark.asyncio + async def test_detects_bot_via_type_field( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + mock_client = mock_github_http_client._client + response = MagicMock() + response.json.return_value = {"login": "forge-bot", "type": "Bot"} + response.raise_for_status = MagicMock() + mock_client.get = AsyncMock(return_value=response) + + actor = await github_adapter_with_mock_client.get_authenticated_identity(github_repo_ref) + + assert actor.is_bot is True + + @pytest.mark.asyncio + async def test_detects_bot_via_login_suffix( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + mock_client = mock_github_http_client._client + response = MagicMock() + response.json.return_value = {"login": "dependabot[bot]", "type": "User"} + response.raise_for_status = MagicMock() + mock_client.get = AsyncMock(return_value=response) + + actor = await github_adapter_with_mock_client.get_authenticated_identity(github_repo_ref) + + assert actor.is_bot is True + + +class TestEnsureWriteTarget: + @pytest.mark.asyncio + async def test_fork_mode_creates_and_syncs_fork( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """Fork mode creates/reuses a fork, syncs it, and derives the target + coordinates from the fork API response.""" + mock_github_http_client.get_or_create_fork = AsyncMock( + return_value={ + "name": "repo", + "owner": {"login": "forge-bot"}, + "clone_url": "https://github.com/forge-bot/repo.git", + "default_branch": "main", + } + ) + mock_github_http_client.sync_fork_with_upstream = AsyncMock(return_value=True) + + target = await github_adapter_with_mock_client.ensure_write_target(github_repo_ref) + + mock_github_http_client.get_or_create_fork.assert_awaited_once_with("test", "repo") + mock_github_http_client.sync_fork_with_upstream.assert_awaited_once_with( + "forge-bot", "repo", branch="main" + ) + assert target.clone_url == "https://github.com/forge-bot/repo.git" + assert target.push_remote_name == "origin" + assert target.head_ref == "forge/test/repo" + assert target.base_branch == "main" + + @pytest.mark.asyncio + async def test_fork_mode_raises_conflict_when_fork_diverged( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """A diverged fork (sync returns False) surfaces as ConflictError.""" + mock_github_http_client.get_or_create_fork = AsyncMock( + return_value={ + "name": "repo", + "owner": {"login": "forge-bot"}, + "clone_url": "https://github.com/forge-bot/repo.git", + } + ) + mock_github_http_client.sync_fork_with_upstream = AsyncMock(return_value=False) + + with pytest.raises(ConflictError, match="diverged"): + await github_adapter_with_mock_client.ensure_write_target(github_repo_ref) + + @pytest.mark.asyncio + async def test_direct_mode_makes_no_api_calls( + self, + github_adapter_with_mock_client: GitHubAdapter, + mock_github_http_client: GitHubClient, + ): + """Direct mode targets upstream directly and touches no HTTP/client methods.""" + direct_repo_ref = RepositoryRef( + id="test/repo", + provider=Provider.GITHUB, + connection="test-github", + namespace="test/repo", + default_branch="develop", + change_request_mode="direct", + ) + + target = await github_adapter_with_mock_client.ensure_write_target(direct_repo_ref) + + assert target.clone_url == "https://github.com/test/repo.git" + assert target.push_remote_name == "origin" + assert target.head_ref == "forge/test/repo" + assert target.base_branch == "develop" + # No fork/sync/HTTP work should have happened in direct mode. + mock_github_http_client._client.get.assert_not_called() + mock_github_http_client._client.post.assert_not_called() + + @pytest.mark.asyncio + async def test_direct_mode_derives_clone_url_from_enterprise_connection(self): + """A GitHub Enterprise Server connection's web host (no 'api.' subdomain, + no '/api/v3' suffix) must be used for the clone URL, not github.com.""" + enterprise_connection = Connection( + name="ghe", + provider=Provider.GITHUB, + base_url="https://ghe.example.com/api/v3", + credential_env="GHE_TOKEN", + webhook_secret_env="GHE_WEBHOOK_SECRET", + ) + adapter = GitHubAdapter(enterprise_connection, credential="test-token-123") + direct_repo_ref = RepositoryRef( + id="test/repo", + provider=Provider.GITHUB, + connection="ghe", + namespace="test/repo", + default_branch="main", + change_request_mode="direct", + ) + + target = await adapter.ensure_write_target(direct_repo_ref) + + assert target.clone_url == "https://ghe.example.com/test/repo.git" + + +class TestLazyClientConnectionPlumbing: + def test_lazy_client_uses_connection_base_url_and_ca_path(self): + """GitHubAdapter's lazily-constructed client must target the configured + connection's API host and CA bundle, not the public GitHub defaults.""" + enterprise_connection = Connection( + name="ghe", + provider=Provider.GITHUB, + base_url="https://ghe.example.com/api/v3", + credential_env="GHE_TOKEN", + webhook_secret_env="GHE_WEBHOOK_SECRET", + ca_path="/etc/ssl/certs/ghe-ca.pem", + ) + adapter = GitHubAdapter(enterprise_connection, credential="test-token-123") + + client = adapter._get_client() + + assert client.base_url == "https://ghe.example.com/api/v3" + assert client._ca_path == "/etc/ssl/certs/ghe-ca.pem" + + @pytest.mark.asyncio + async def test_close_closes_the_lazily_constructed_client(self, github_connection: Connection): + adapter = GitHubAdapter(github_connection, credential="test-token-123") + client = adapter._get_client() + client.close = AsyncMock() + + await adapter.close() + + client.close.assert_awaited_once() + + @pytest.mark.asyncio + async def test_close_is_a_noop_when_no_client_was_ever_constructed( + self, github_connection: Connection + ): + adapter = GitHubAdapter(github_connection, credential="test-token-123") + + await adapter.close() # must not raise, must not construct a client + + assert adapter._client is None + + +class TestGetGitCredentials: + @pytest.mark.asyncio + async def test_public_github_derives_bare_host_and_credential( + self, github_adapter: GitHubAdapter, github_repo_ref: RepositoryRef + ): + """Public GitHub's web host (github.com) must be used for git + operations, not the API host (api.github.com).""" + credentials = await github_adapter.get_git_credentials(github_repo_ref) + + assert credentials.host == "github.com" + assert credentials.token == "test-token-123" + assert credentials.ca_path is None + + @pytest.mark.asyncio + async def test_enterprise_connection_derives_web_host_and_ca_path( + self, github_repo_ref: RepositoryRef + ): + """An Enterprise Server connection's git host has no /api/v3 suffix + (unlike its API base_url), and its CA bundle must be carried through + for git's own TLS verification.""" + enterprise_connection = Connection( + name="ghe", + provider=Provider.GITHUB, + base_url="https://ghe.example.com/api/v3", + credential_env="GHE_TOKEN", + webhook_secret_env="GHE_WEBHOOK_SECRET", + ca_path="/etc/ssl/certs/ghe-ca.pem", + ) + adapter = GitHubAdapter(enterprise_connection, credential="ghe-token-456") + + credentials = await adapter.get_git_credentials(github_repo_ref) + + assert credentials.host == "ghe.example.com" + assert credentials.token == "ghe-token-456" + assert credentials.ca_path == "/etc/ssl/certs/ghe-ca.pem" + + +@pytest.fixture +def write_target() -> WriteTarget: + return WriteTarget( + clone_url="https://github.com/forge-bot/repo.git", + push_remote_name="origin", + head_ref="forge/test/repo", + base_branch="main", + ) + + +def _pr_dict(**overrides) -> dict: + """A representative GitHub PR API response, with optional field overrides.""" + pr = { + "number": 42, + "html_url": "https://github.com/test/repo/pull/42", + "title": "Test PR", + "body": "Test body", + "state": "open", + "merged": False, + "draft": False, + "head": {"ref": "forge/test/repo"}, + "base": {"ref": "main"}, + } + pr.update(overrides) + return pr + + +class TestCreateChangeRequest: + @pytest.mark.asyncio + async def test_creates_pr_and_maps_result( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + write_target: WriteTarget, + ): + mock_github_http_client.create_pull_request = AsyncMock( + return_value=PullRequestCreationResult(pr=_pr_dict(), created=True) + ) + + cr = await github_adapter_with_mock_client.create_change_request( + github_repo_ref, + write_target, + title="Test PR", + body="Test body", + draft=False, + ) + + mock_github_http_client.create_pull_request.assert_awaited_once_with( + owner="test", + repo="repo", + title="Test PR", + body="Test body", + head="forge/test/repo", + base="main", + draft=False, + ) + assert cr.identity == ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=42 + ) + assert cr.url == "https://github.com/test/repo/pull/42" + assert cr.title == "Test PR" + assert cr.body == "Test body" + assert cr.state == ChangeRequestState.OPEN + assert cr.source_branch == "forge/test/repo" + assert cr.target_branch == "main" + assert cr.draft is False + + @pytest.mark.asyncio + async def test_returns_existing_pr_when_already_present( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + write_target: WriteTarget, + ): + """The 422-already-exists path returns created=False with the existing PR, + which is mapped and returned like any other PR.""" + mock_github_http_client.create_pull_request = AsyncMock( + return_value=PullRequestCreationResult( + pr=_pr_dict(number=7, html_url="https://github.com/test/repo/pull/7"), + created=False, + ) + ) + + cr = await github_adapter_with_mock_client.create_change_request( + github_repo_ref, + write_target, + title="Test PR", + body="Test body", + ) + + assert cr.identity.native_id == 7 + assert cr.url == "https://github.com/test/repo/pull/7" + assert cr.state == ChangeRequestState.OPEN + + @pytest.mark.asyncio + async def test_draft_flag_is_forwarded( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + write_target: WriteTarget, + ): + mock_github_http_client.create_pull_request = AsyncMock( + return_value=PullRequestCreationResult(pr=_pr_dict(draft=True), created=True) + ) + + cr = await github_adapter_with_mock_client.create_change_request( + github_repo_ref, + write_target, + title="Test PR", + body="Test body", + draft=True, + ) + + assert mock_github_http_client.create_pull_request.await_args.kwargs["draft"] is True + assert cr.draft is True + + +def test_map_change_request_rejects_both_repo_ref_and_identity( + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, +): + """_map_change_request's contract is "exactly one of repo_ref or identity" -- + passing both must raise rather than silently letting identity win.""" + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=1 + ) + + with pytest.raises(ValueError, match="not both"): + github_adapter_with_mock_client._map_change_request( + _pr_dict(), repo_ref=github_repo_ref, identity=identity + ) + + +class TestGetChangeRequest: + @pytest.mark.asyncio + async def test_fetches_and_preserves_identity( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=42 + ) + mock_github_http_client.get_pull_request = AsyncMock( + return_value=_pr_dict(state="closed", merged=True) + ) + + cr = await github_adapter_with_mock_client.get_change_request(github_repo_ref, identity) + + mock_github_http_client.get_pull_request.assert_awaited_once_with("test", "repo", 42) + # The exact identity object passed in is preserved, not reconstructed. + assert cr.identity is identity + assert cr.state == ChangeRequestState.MERGED + + @pytest.mark.asyncio + async def test_coerces_string_native_id_to_int( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id="99" + ) + mock_github_http_client.get_pull_request = AsyncMock(return_value=_pr_dict(number=99)) + + await github_adapter_with_mock_client.get_change_request(github_repo_ref, identity) + + mock_github_http_client.get_pull_request.assert_awaited_once_with("test", "repo", 99) + + @pytest.mark.asyncio + async def test_raises_on_missing_native_id_instead_of_defaulting_to_zero( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=None + ) + mock_github_http_client.get_pull_request = AsyncMock() + + with pytest.raises(ValueError, match="native_id"): + await github_adapter_with_mock_client.get_change_request(github_repo_ref, identity) + + mock_github_http_client.get_pull_request.assert_not_awaited() + + +class TestUpdateChangeRequest: + @pytest.mark.asyncio + async def test_updates_title_and_body( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=42 + ) + mock_github_http_client.update_pull_request = AsyncMock( + return_value=_pr_dict(title="New title", body="New body") + ) + + cr = await github_adapter_with_mock_client.update_change_request( + github_repo_ref, + identity, + title="New title", + body="New body", + ) + + mock_github_http_client.update_pull_request.assert_awaited_once_with( + owner="test", + repo="repo", + pr_number=42, + title="New title", + body="New body", + state=None, + ) + assert cr.identity is identity + assert cr.title == "New title" + assert cr.body == "New body" + + @pytest.mark.asyncio + async def test_maps_closed_state_to_github_string( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=42 + ) + mock_github_http_client.update_pull_request = AsyncMock( + return_value=_pr_dict(state="closed") + ) + + cr = await github_adapter_with_mock_client.update_change_request( + github_repo_ref, + identity, + state=ChangeRequestState.CLOSED, + ) + + assert mock_github_http_client.update_pull_request.await_args.kwargs["state"] == "closed" + assert cr.state == ChangeRequestState.CLOSED + + @pytest.mark.asyncio + async def test_maps_open_state_to_github_string( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=42 + ) + mock_github_http_client.update_pull_request = AsyncMock(return_value=_pr_dict(state="open")) + + await github_adapter_with_mock_client.update_change_request( + github_repo_ref, + identity, + state=ChangeRequestState.OPEN, + ) + + assert mock_github_http_client.update_pull_request.await_args.kwargs["state"] == "open" + + @pytest.mark.asyncio + async def test_rejects_merged_state( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """MERGED cannot be set via the update endpoint; the adapter raises rather + than silently ignoring it or closing the PR instead.""" + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=42 + ) + mock_github_http_client.update_pull_request = AsyncMock() + + with pytest.raises(ValueError, match="MERGED"): + await github_adapter_with_mock_client.update_change_request( + github_repo_ref, + identity, + state=ChangeRequestState.MERGED, + ) + + # No API call should have been made when the state is rejected. + mock_github_http_client.update_pull_request.assert_not_awaited() + + @pytest.mark.asyncio + async def test_raises_on_missing_native_id_instead_of_defaulting_to_zero( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """_require_native_id's guard applies here too, not just get_change_request.""" + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=None + ) + mock_github_http_client.update_pull_request = AsyncMock() + + with pytest.raises(ValueError, match="native_id"): + await github_adapter_with_mock_client.update_change_request( + github_repo_ref, identity, title="New title" + ) + + mock_github_http_client.update_pull_request.assert_not_awaited() + + +class TestCreateComment: + @pytest.mark.asyncio + async def test_creates_issue_comment_and_maps_result( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=42 + ) + mock_github_http_client.create_issue_comment = AsyncMock( + return_value={ + "id": 555, + "body": "Looks good", + "user": {"login": "forge-bot"}, + } + ) + + comment = await github_adapter_with_mock_client.create_comment( + github_repo_ref, identity, "Looks good" + ) + + mock_github_http_client.create_issue_comment.assert_awaited_once_with( + owner="test", + repo="repo", + issue_number=42, + body="Looks good", + ) + assert comment.id == "555" + assert comment.body == "Looks good" + assert comment.author == "forge-bot" + assert comment.in_reply_to is None + + @pytest.mark.asyncio + async def test_coerces_string_native_id_to_int( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id="99" + ) + mock_github_http_client.create_issue_comment = AsyncMock( + return_value={"id": 1, "body": "x", "user": {"login": "forge-bot"}} + ) + + await github_adapter_with_mock_client.create_comment(github_repo_ref, identity, "x") + + assert mock_github_http_client.create_issue_comment.await_args.kwargs["issue_number"] == 99 + + @pytest.mark.asyncio + async def test_handles_null_user_from_deleted_account( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """GitHub returns "user": null for comments from a deleted account — + this must not crash with AttributeError on None.get(...).""" + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=42 + ) + mock_github_http_client.create_issue_comment = AsyncMock( + return_value={"id": 1, "body": "x", "user": None} + ) + + comment = await github_adapter_with_mock_client.create_comment( + github_repo_ref, identity, "x" + ) + + assert comment.author == "" + + +class TestReplyToComment: + @pytest.mark.asyncio + async def test_uses_reply_endpoint_and_sets_in_reply_to( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=42 + ) + mock_github_http_client.reply_to_review_comment = AsyncMock( + return_value={ + "id": 777, + "body": "Thanks", + "user": {"login": "forge-bot"}, + "path": "src/foo.py", + "line": 10, + } + ) + # Ensure the generic issue-comment endpoint is NOT used for replies. + mock_github_http_client.create_issue_comment = AsyncMock() + + comment = await github_adapter_with_mock_client.reply_to_comment( + github_repo_ref, identity, comment_id="123", body="Thanks" + ) + + mock_github_http_client.reply_to_review_comment.assert_awaited_once_with( + owner="test", + repo="repo", + pr_number=42, + comment_id=123, + body="Thanks", + ) + mock_github_http_client.create_issue_comment.assert_not_awaited() + assert comment.id == "777" + assert comment.body == "Thanks" + assert comment.author == "forge-bot" + assert comment.path == "src/foo.py" + assert comment.line == 10 + assert comment.in_reply_to == "123" + + +class TestGetReviewThreads: + @pytest.mark.asyncio + async def test_maps_approved_and_changes_requested( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=42 + ) + mock_github_http_client.get_reviews = AsyncMock( + return_value=[ + {"id": 1, "state": "APPROVED", "body": "LGTM", "user": {"login": "alice"}}, + {"id": 2, "state": "CHANGES_REQUESTED", "body": "Fix", "user": {"login": "bob"}}, + ] + ) + + reviews = await github_adapter_with_mock_client.get_review_threads( + github_repo_ref, identity + ) + + mock_github_http_client.get_reviews.assert_awaited_once_with("test", "repo", 42) + assert [r.id for r in reviews] == ["1", "2"] + assert reviews[0].state == ReviewState.APPROVED + assert reviews[0].body == "LGTM" + assert reviews[0].author == "alice" + assert reviews[0].comments == [] + assert reviews[1].state == ReviewState.CHANGES_REQUESTED + assert reviews[1].author == "bob" + assert reviews[1].comments == [] + + @pytest.mark.asyncio + async def test_handles_null_user_from_deleted_account( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """GitHub returns "user": null for a review from a deleted account — + this must not crash with AttributeError on None.get(...).""" + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=42 + ) + mock_github_http_client.get_reviews = AsyncMock( + return_value=[{"id": 3, "state": "COMMENTED", "body": None, "user": None}] + ) + + reviews = await github_adapter_with_mock_client.get_review_threads( + github_repo_ref, identity + ) + + assert reviews[0].author == "" + assert reviews[0].body == "" + assert reviews[0].state == ReviewState.COMMENTED + + @pytest.mark.asyncio + async def test_dismissed_maps_to_dismissed( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """DISMISSED has its own dedicated ReviewState member, kept distinct + from COMMENTED so a withdrawn review isn't mistaken for an active one + requesting changes (an admin dismissing a stale review must not + trigger a revision request downstream).""" + identity = ChangeRequestIdentity( + connection="test-github", repository_id="test/repo", native_id=42 + ) + mock_github_http_client.get_reviews = AsyncMock( + return_value=[ + {"id": 4, "state": "DISMISSED", "body": "old", "user": {"login": "carol"}}, + {"id": 5, "state": "PENDING", "body": "", "user": {"login": "dave"}}, + ] + ) + + reviews = await github_adapter_with_mock_client.get_review_threads( + github_repo_ref, identity + ) + + assert reviews[0].state == ReviewState.DISMISSED + assert reviews[1].state == ReviewState.PENDING + + +class TestGetChecks: + @pytest.mark.asyncio + async def test_maps_actions_backed_check_run_with_logs_url( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """An Actions-backed check run stores the workflow *run* id (parsed from + details_url) in logs_url -- deliberately distinct from the check-run id, + since a check-run id is not an Actions job id.""" + mock_github_http_client.get_check_runs = AsyncMock( + return_value=[ + { + "id": 12345, # check-run id: NOT the run id, NOT the job id + "name": "build", + "status": "completed", + "conclusion": "success", + "html_url": "https://github.com/test/repo/runs/12345", + "details_url": "https://github.com/test/repo/actions/runs/987654", + "app": {"slug": "github-actions"}, + } + ] + ) + + checks = await github_adapter_with_mock_client.get_checks(github_repo_ref, "abc123") + + mock_github_http_client.get_check_runs.assert_awaited_once_with( + owner="test", repo="repo", ref="abc123" + ) + assert len(checks) == 1 + check = checks[0] + assert check.name == "build" + assert check.status == CheckStatus.COMPLETED + assert check.conclusion == CheckConclusion.SUCCESS + assert check.url == "https://github.com/test/repo/runs/12345" + # The run id from details_url, not the check-run id. + assert check.logs_url == "987654" + + @pytest.mark.asyncio + async def test_non_actions_app_with_numeric_id_has_no_logs_url( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """A check run from a non-Actions GitHub App (e.g. CodeQL) is not backed + by an Actions job -- even though it carries a numeric id and an + actions-style details_url, logs_url must stay None rather than pointing + get_check_logs at a run that has no matching Actions job.""" + mock_github_http_client.get_check_runs = AsyncMock( + return_value=[ + { + "id": 99999, + "name": "CodeQL", + "status": "completed", + "conclusion": "success", + "html_url": "https://github.com/test/repo/runs/99999", + "details_url": "https://github.com/test/repo/actions/runs/99999", + "app": {"slug": "github-code-scanning"}, + } + ] + ) + + checks = await github_adapter_with_mock_client.get_checks(github_repo_ref, "abc123") + + assert checks[0].logs_url is None + + @pytest.mark.asyncio + async def test_actions_check_without_run_id_in_details_url_has_no_logs_url( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """An Actions check whose details_url isn't the expected + /actions/runs/{id} shape yields no resolvable run id, so logs_url stays + None rather than storing a bogus resolution key.""" + mock_github_http_client.get_check_runs = AsyncMock( + return_value=[ + { + "id": 12345, + "name": "build", + "status": "completed", + "conclusion": "success", + "details_url": "https://example.com/some/other/path", + "app": {"slug": "github-actions"}, + } + ] + ) + + checks = await github_adapter_with_mock_client.get_checks(github_repo_ref, "abc123") + + assert checks[0].logs_url is None + + @pytest.mark.asyncio + async def test_timed_out_and_action_required_map_to_failure( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """timed_out and action_required both indicate the check did not pass, + so they must not collapse into the same NONE bucket as "no conclusion + yet".""" + mock_github_http_client.get_check_runs = AsyncMock( + return_value=[ + {"id": 1, "name": "slow-job", "status": "completed", "conclusion": "timed_out"}, + {"id": 2, "name": "gate", "status": "completed", "conclusion": "action_required"}, + ] + ) + + checks = await github_adapter_with_mock_client.get_checks(github_repo_ref, "abc123") + + assert checks[0].conclusion == CheckConclusion.FAILURE + assert checks[1].conclusion == CheckConclusion.FAILURE + + @pytest.mark.asyncio + async def test_maps_commit_status_backed_entry_with_no_logs_url( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """A commit-status-backed check (per _normalize_commit_status) has no id, + so logs_url must be None — there is no logs endpoint for it.""" + mock_github_http_client.get_check_runs = AsyncMock( + return_value=[ + { + "name": "prow/verify", + "status": "completed", + "conclusion": "failure", + "output": {"summary": "it broke"}, + "html_url": "https://prow.example.com/log", + } + ] + ) + + checks = await github_adapter_with_mock_client.get_checks(github_repo_ref, "abc123") + + assert len(checks) == 1 + check = checks[0] + assert check.name == "prow/verify" + assert check.status == CheckStatus.COMPLETED + assert check.conclusion == CheckConclusion.FAILURE + assert check.url == "https://prow.example.com/log" + assert check.logs_url is None + + @pytest.mark.asyncio + async def test_unknown_conclusion_maps_to_none( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """A missing/unrecognized conclusion becomes CheckConclusion.NONE, and a + missing status with no conclusion is inferred as IN_PROGRESS.""" + mock_github_http_client.get_check_runs = AsyncMock( + return_value=[ + {"id": 1, "name": "pending-check", "status": "", "conclusion": None}, + ] + ) + + checks = await github_adapter_with_mock_client.get_checks(github_repo_ref, "abc123") + + assert checks[0].conclusion == CheckConclusion.NONE + assert checks[0].status == CheckStatus.IN_PROGRESS + assert checks[0].url == "" + + +class TestGetCheckLogs: + @pytest.mark.asyncio + async def test_fetches_logs_for_actions_backed_check( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """logs_url holds the workflow *run* id; the adapter lists the run's + jobs, matches by name, and fetches the *job* id's logs. Run id, job id, + and the original check-run id are all deliberately distinct so the test + can't pass by treating any of them as interchangeable.""" + mock_github_http_client.list_workflow_run_jobs = AsyncMock( + return_value=[ + {"id": 555, "name": "lint"}, + {"id": 777, "name": "build"}, + ] + ) + mock_github_http_client.get_job_logs = AsyncMock(return_value="line 1\nline 2\n") + check = CheckRun( + name="build", + status=CheckStatus.COMPLETED, + conclusion=CheckConclusion.SUCCESS, + url="https://github.com/test/repo/runs/12345", + logs_url="987654", # workflow run id, not the check-run or job id + ) + + logs = await github_adapter_with_mock_client.get_check_logs(github_repo_ref, check) + + mock_github_http_client.list_workflow_run_jobs.assert_awaited_once_with( + "test", "repo", 987654 + ) + # The resolved job id (777), NOT the run id (987654) or check-run id. + mock_github_http_client.get_job_logs.assert_awaited_once_with("test", "repo", job_id=777) + assert logs == "line 1\nline 2\n" + + @pytest.mark.asyncio + async def test_raises_not_found_when_no_logs_url( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """A commit-status check (logs_url None) has no logs endpoint, so + requesting its logs raises NotFoundError without any HTTP call.""" + mock_github_http_client.list_workflow_run_jobs = AsyncMock() + mock_github_http_client.get_job_logs = AsyncMock() + check = CheckRun( + name="prow/verify", + status=CheckStatus.COMPLETED, + conclusion=CheckConclusion.FAILURE, + url="https://prow.example.com/log", + logs_url=None, + ) + + with pytest.raises(NotFoundError, match="No logs available"): + await github_adapter_with_mock_client.get_check_logs(github_repo_ref, check) + + mock_github_http_client.list_workflow_run_jobs.assert_not_awaited() + mock_github_http_client.get_job_logs.assert_not_awaited() + + @pytest.mark.asyncio + async def test_raises_source_control_error_for_non_numeric_logs_url( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """logs_url survives a round-trip through the Redis queue as a raw string; + a malformed/legacy value must raise the documented SourceControlError + rather than leaking an uncaught ValueError, and must not hit the API.""" + mock_github_http_client.list_workflow_run_jobs = AsyncMock() + check = CheckRun( + name="build", + status=CheckStatus.COMPLETED, + conclusion=CheckConclusion.SUCCESS, + logs_url="not-a-number", + ) + + with pytest.raises(SourceControlError, match="non-numeric logs_url"): + await github_adapter_with_mock_client.get_check_logs(github_repo_ref, check) + + mock_github_http_client.list_workflow_run_jobs.assert_not_awaited() + + @pytest.mark.asyncio + async def test_raises_not_found_when_no_job_matches_check_name( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """If the run has no job whose name matches the check, there is no job id + to fetch logs for -- raise NotFoundError rather than guessing.""" + mock_github_http_client.list_workflow_run_jobs = AsyncMock( + return_value=[ + {"id": 555, "name": "lint"}, + {"id": 666, "name": "test"}, + ] + ) + mock_github_http_client.get_job_logs = AsyncMock() + check = CheckRun( + name="build", + status=CheckStatus.COMPLETED, + conclusion=CheckConclusion.SUCCESS, + logs_url="987654", + ) + + with pytest.raises(NotFoundError, match="No Actions job named 'build'"): + await github_adapter_with_mock_client.get_check_logs(github_repo_ref, check) + + mock_github_http_client.get_job_logs.assert_not_awaited() + + @pytest.mark.asyncio + async def test_raises_when_multiple_jobs_match_check_name( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """If more than one job in the run shares the check's name, the correct + logs can't be chosen unambiguously -- raise rather than fetch the wrong + job's logs.""" + mock_github_http_client.list_workflow_run_jobs = AsyncMock( + return_value=[ + {"id": 777, "name": "build"}, + {"id": 888, "name": "build"}, + ] + ) + mock_github_http_client.get_job_logs = AsyncMock() + check = CheckRun( + name="build", + status=CheckStatus.COMPLETED, + conclusion=CheckConclusion.SUCCESS, + logs_url="987654", + ) + + with pytest.raises(SourceControlError, match="2 jobs named 'build'"): + await github_adapter_with_mock_client.get_check_logs(github_repo_ref, check) + + mock_github_http_client.get_job_logs.assert_not_awaited() + + @pytest.mark.asyncio + async def test_translates_404_into_not_found( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """A 404 httpx.HTTPStatusError from get_job_logs is translated into the + provider-neutral NotFoundError.""" + mock_github_http_client.list_workflow_run_jobs = AsyncMock( + return_value=[{"id": 777, "name": "build"}] + ) + response = httpx.Response(404, request=httpx.Request("GET", "https://api.github.com/logs")) + mock_github_http_client.get_job_logs = AsyncMock( + side_effect=httpx.HTTPStatusError( + "not found", request=response.request, response=response + ) + ) + check = CheckRun( + name="build", + status=CheckStatus.COMPLETED, + conclusion=CheckConclusion.SUCCESS, + url="https://github.com/test/repo/runs/12345", + logs_url="987654", + ) + + with pytest.raises(NotFoundError, match="were not"): + await github_adapter_with_mock_client.get_check_logs(github_repo_ref, check) + + @pytest.mark.asyncio + async def test_non_404_http_error_propagates( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """Non-404 HTTP errors are not swallowed as NotFoundError; they are + translated into the neutral error hierarchy at the adapter boundary + instead of leaking a raw httpx exception.""" + mock_github_http_client.list_workflow_run_jobs = AsyncMock( + return_value=[{"id": 777, "name": "build"}] + ) + response = httpx.Response(500, request=httpx.Request("GET", "https://api.github.com/logs")) + mock_github_http_client.get_job_logs = AsyncMock( + side_effect=httpx.HTTPStatusError("boom", request=response.request, response=response) + ) + check = CheckRun( + name="build", + status=CheckStatus.COMPLETED, + conclusion=CheckConclusion.SUCCESS, + logs_url="987654", + ) + + with pytest.raises(TransientProviderError): + await github_adapter_with_mock_client.get_check_logs(github_repo_ref, check) + + +@pytest.mark.asyncio +async def test_repository_operations( + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, +): + """Conformance suite coverage for repository metadata operations. + + Exercises resolve_default_branch and get_authenticated_identity back-to-back + on the same adapter/client instance, which also guards against per-call + client teardown silently swapping the mocked httpx.AsyncClient for a real one. + """ + mock_client = mock_github_http_client._client + + repo_response = MagicMock() + repo_response.json.return_value = {"default_branch": "main", "full_name": "test/repo"} + repo_response.raise_for_status = MagicMock() + + user_response = MagicMock() + user_response.json.return_value = {"login": "octocat", "type": "User"} + user_response.raise_for_status = MagicMock() + + mock_client.get = AsyncMock(side_effect=[repo_response, user_response]) + + await assert_repository_operations(github_adapter_with_mock_client, github_repo_ref) + + +class TestGetFile: + @pytest.mark.asyncio + async def test_returns_decoded_content( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """A found file's base64 content (with the embedded newlines GitHub's API + wraps it in) round-trips back to the original UTF-8 text.""" + original = "line one\nline two\nline three\n" * 5 + # GitHub wraps base64 content at column 76 with embedded newlines. + encoded = base64.encodebytes(original.encode()).decode() + assert "\n" in encoded.strip() # sanity: multi-line payload + mock_github_http_client.get_file_contents = AsyncMock( + return_value={ + "content": encoded, + "encoding": "base64", + "sha": "abc123", + "path": "README.md", + } + ) + + content = await github_adapter_with_mock_client.get_file( + github_repo_ref, "README.md", "main" + ) + + mock_github_http_client.get_file_contents.assert_awaited_once_with( + owner="test", repo="repo", path="README.md", ref="main" + ) + assert content == original + + @pytest.mark.asyncio + async def test_raises_not_found_when_missing( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """A missing file (client returns None for 404) surfaces as NotFoundError, + not a silent empty string, because the protocol return type is str.""" + mock_github_http_client.get_file_contents = AsyncMock(return_value=None) + + with pytest.raises(NotFoundError, match="not found"): + await github_adapter_with_mock_client.get_file(github_repo_ref, "missing.txt", "main") + + @pytest.mark.asyncio + async def test_raises_on_file_too_large_for_inline_content( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """GitHub only inlines base64 content for files <=1MB; larger files come + back with encoding "none" and empty content on an otherwise-successful + response. That must raise, not silently look like an empty file.""" + mock_github_http_client.get_file_contents = AsyncMock( + return_value={"content": "", "encoding": "none", "size": 5_000_000, "sha": "big-sha"} + ) + + with pytest.raises(SourceControlError, match="1MB"): + await github_adapter_with_mock_client.get_file(github_repo_ref, "big.bin", "main") + + +class TestPutFile: + @pytest.mark.asyncio + async def test_creates_new_file_without_sha( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """When the file doesn't exist yet (lookup returns None), no sha is passed + so the Contents API treats it as a create.""" + mock_github_http_client.get_file_contents = AsyncMock(return_value=None) + mock_github_http_client.create_or_update_file = AsyncMock(return_value={}) + + await github_adapter_with_mock_client.put_file( + github_repo_ref, + path="docs/new.md", + content="hello", + message="add new.md", + branch="feature", + ) + + mock_github_http_client.get_file_contents.assert_awaited_once_with( + owner="test", repo="repo", path="docs/new.md", ref="feature" + ) + mock_github_http_client.create_or_update_file.assert_awaited_once_with( + owner="test", + repo="repo", + path="docs/new.md", + content="hello", + message="add new.md", + branch="feature", + sha=None, + ) + + @pytest.mark.asyncio + async def test_updates_existing_file_passes_sha_through( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """When the file already exists, its sha from the prior lookup is threaded + into the update call (a create-without-sha would 422).""" + mock_github_http_client.get_file_contents = AsyncMock( + return_value={"content": "b2xk", "sha": "existing-sha-999"} + ) + mock_github_http_client.create_or_update_file = AsyncMock(return_value={}) + + await github_adapter_with_mock_client.put_file( + github_repo_ref, + path="docs/existing.md", + content="updated body", + message="update existing.md", + branch="main", + ) + + mock_github_http_client.get_file_contents.assert_awaited_once_with( + owner="test", repo="repo", path="docs/existing.md", ref="main" + ) + mock_github_http_client.create_or_update_file.assert_awaited_once_with( + owner="test", + repo="repo", + path="docs/existing.md", + content="updated body", + message="update existing.md", + branch="main", + sha="existing-sha-999", + ) + + @pytest.mark.asyncio + async def test_translates_stale_sha_conflict( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """A concurrent write between the sha lookup and the update surfaces as + a 409/422 from GitHub; that must translate into ConflictError rather + than propagate as a raw httpx error.""" + mock_github_http_client.get_file_contents = AsyncMock( + return_value={"content": "b2xk", "sha": "stale-sha"} + ) + response = httpx.Response( + 409, request=httpx.Request("PUT", "https://api.github.com/contents") + ) + mock_github_http_client.create_or_update_file = AsyncMock( + side_effect=httpx.HTTPStatusError( + "conflict", request=response.request, response=response + ) + ) + + with pytest.raises(ConflictError, match="concurrently modified"): + await github_adapter_with_mock_client.put_file( + github_repo_ref, + path="docs/existing.md", + content="updated body", + message="update existing.md", + branch="main", + ) + + @pytest.mark.asyncio + async def test_translates_422_with_sha_message_to_conflict( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """A 422 whose error message mentions the sha mismatch is the same + stale-sha conflict as a 409 -- must still translate to ConflictError.""" + mock_github_http_client.get_file_contents = AsyncMock( + return_value={"content": "b2xk", "sha": "stale-sha"} + ) + response = httpx.Response( + 422, + json={"message": "docs/existing.md does not match stale-sha"}, + request=httpx.Request("PUT", "https://api.github.com/contents"), + ) + mock_github_http_client.create_or_update_file = AsyncMock( + side_effect=httpx.HTTPStatusError( + "unprocessable", request=response.request, response=response + ) + ) + + with pytest.raises(ConflictError, match="concurrently modified"): + await github_adapter_with_mock_client.put_file( + github_repo_ref, + path="docs/existing.md", + content="updated body", + message="update existing.md", + branch="main", + ) + + @pytest.mark.asyncio + async def test_unrelated_422_propagates_without_conflict_mislabel( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """A 422 for a reason unrelated to a stale sha (e.g. branch protection) + must not be mislabeled as a retryable ConflictError.""" + mock_github_http_client.get_file_contents = AsyncMock( + return_value={"content": "b2xk", "sha": "current-sha"} + ) + response = httpx.Response( + 422, + json={"message": "Changes must be made through a pull request"}, + request=httpx.Request("PUT", "https://api.github.com/contents"), + ) + error = httpx.HTTPStatusError("unprocessable", request=response.request, response=response) + mock_github_http_client.create_or_update_file = AsyncMock(side_effect=error) + + with pytest.raises(httpx.HTTPStatusError): + await github_adapter_with_mock_client.put_file( + github_repo_ref, + path="docs/existing.md", + content="updated body", + message="update existing.md", + branch="main", + ) + + +class TestParseWebhookEnrichment: + @pytest.mark.asyncio + async def test_issue_comment_on_pr_populates_change_request_and_comment( + self, + github_adapter: GitHubAdapter, + github_repo_ref: RepositoryRef, + github_connection: Connection, + ): + """An issue_comment event on a PR carries the PR's identity under + payload["issue"]["pull_request"], not top-level payload["pull_request"] — + this must still populate change_request, plus the comment field.""" + payload = { + "action": "created", + "issue": { + "number": 42, + "title": "Add feature", + "body": "PR description", + "state": "open", + "html_url": "https://github.com/test/repo/issues/42", + "pull_request": {"url": "https://api.github.com/repos/test/repo/pulls/42"}, + }, + "comment": { + "id": 555, + "body": "Looks good", + "user": {"login": "reviewer1"}, + }, + "repository": {"full_name": "test/repo"}, + "sender": {"login": "reviewer1", "type": "User"}, + } + body = json.dumps(payload).encode() + headers = {"X-GitHub-Event": "issue_comment", "X-GitHub-Delivery": "id-1"} + resolver = MockResolver(github_repo_ref, github_connection) + + event = await github_adapter.parse_webhook(headers, body, resolver) + + assert event.kind == EventKind.COMMENT_CREATED + assert event.change_request is not None + assert event.change_request.identity.native_id == 42 + assert event.change_request.title == "Add feature" + assert event.comment is not None + assert event.comment.id == "555" + assert event.comment.body == "Looks good" + assert event.comment.author == "reviewer1" + + @pytest.mark.asyncio + async def test_issue_comment_on_plain_issue_has_no_change_request( + self, + github_adapter: GitHubAdapter, + github_repo_ref: RepositoryRef, + github_connection: Connection, + ): + """An issue_comment on a plain issue (no "pull_request" key under + "issue") must not fabricate a change_request.""" + payload = { + "action": "created", + "issue": {"number": 7, "title": "Bug report", "state": "open"}, + "comment": {"id": 1, "body": "confirmed", "user": {"login": "x"}}, + "repository": {"full_name": "test/repo"}, + "sender": {"login": "x", "type": "User"}, + } + body = json.dumps(payload).encode() + headers = {"X-GitHub-Event": "issue_comment", "X-GitHub-Delivery": "id-2"} + resolver = MockResolver(github_repo_ref, github_connection) + + event = await github_adapter.parse_webhook(headers, body, resolver) + + assert event.change_request is None + assert event.comment is not None + + @pytest.mark.asyncio + async def test_pull_request_review_comment_populates_comment_with_path_line( + self, + github_adapter: GitHubAdapter, + github_repo_ref: RepositoryRef, + github_connection: Connection, + ): + payload = { + "action": "created", + "comment": { + "id": 999, + "body": "nit: rename this", + "user": {"login": "reviewer2"}, + "path": "src/foo.py", + "line": 42, + "in_reply_to_id": 888, + }, + "pull_request": { + "number": 10, + "html_url": "https://github.com/test/repo/pull/10", + "title": "Fix bug", + "body": "", + "state": "open", + "draft": False, + "head": {"ref": "fix"}, + "base": {"ref": "main"}, + }, + "repository": {"full_name": "test/repo"}, + "sender": {"login": "reviewer2", "type": "User"}, + } + body = json.dumps(payload).encode() + headers = { + "X-GitHub-Event": "pull_request_review_comment", + "X-GitHub-Delivery": "id-3", + } + resolver = MockResolver(github_repo_ref, github_connection) + + event = await github_adapter.parse_webhook(headers, body, resolver) + + assert event.comment is not None + assert event.comment.path == "src/foo.py" + assert event.comment.line == 42 + assert event.comment.in_reply_to == "888" + + @pytest.mark.asyncio + async def test_pull_request_review_populates_review( + self, + github_adapter: GitHubAdapter, + github_repo_ref: RepositoryRef, + github_connection: Connection, + ): + payload = { + "action": "submitted", + "review": { + "id": 321, + "state": "changes_requested", + "body": "please fix", + "user": {"login": "reviewer3"}, + }, + "pull_request": { + "number": 11, + "html_url": "https://github.com/test/repo/pull/11", + "title": "Add thing", + "body": "", + "state": "open", + "draft": False, + "head": {"ref": "add-thing"}, + "base": {"ref": "main"}, + }, + "repository": {"full_name": "test/repo"}, + "sender": {"login": "reviewer3", "type": "User"}, + } + body = json.dumps(payload).encode() + headers = {"X-GitHub-Event": "pull_request_review", "X-GitHub-Delivery": "id-4"} + resolver = MockResolver(github_repo_ref, github_connection) + + event = await github_adapter.parse_webhook(headers, body, resolver) + + assert event.review is not None + assert event.review.id == "321" + assert event.review.state == ReviewState.CHANGES_REQUESTED + assert event.review.author == "reviewer3" + + @pytest.mark.asyncio + async def test_check_run_populates_check( + self, + github_adapter: GitHubAdapter, + github_repo_ref: RepositoryRef, + github_connection: Connection, + ): + payload = { + "action": "completed", + "check_run": { + "id": 654, # check-run id, distinct from the workflow run id + "name": "build", + "status": "completed", + "conclusion": "failure", + "html_url": "https://github.com/test/repo/runs/654", + "details_url": "https://github.com/test/repo/actions/runs/987", + "app": {"slug": "github-actions"}, + "pull_requests": [{"number": 12}], + }, + "repository": {"full_name": "test/repo"}, + "sender": {"login": "github-actions[bot]", "type": "Bot"}, + } + body = json.dumps(payload).encode() + headers = {"X-GitHub-Event": "check_run", "X-GitHub-Delivery": "id-5"} + resolver = MockResolver(github_repo_ref, github_connection) + + event = await github_adapter.parse_webhook(headers, body, resolver) + + assert event.check is not None + assert event.check.name == "build" + assert event.check.conclusion == CheckConclusion.FAILURE + # logs_url holds the workflow run id (987), not the check-run id (654). + assert event.check.logs_url == "987" + # check_run.pull_requests[] is a "simple pull request" stub (number/ + # head/base only) -- without this, a real check_run webhook could + # never be matched to its implementation PR downstream. + assert event.change_request is not None + assert event.change_request.identity.native_id == 12 + + @pytest.mark.asyncio + async def test_check_run_falls_back_to_nested_check_suite_pull_requests( + self, + github_adapter: GitHubAdapter, + github_repo_ref: RepositoryRef, + github_connection: Connection, + ): + """Some check_run payloads carry the PR list only on the nested + check_suite object, not on check_run itself.""" + payload = { + "action": "completed", + "check_run": { + "id": 654, + "name": "build", + "status": "completed", + "conclusion": "success", + "app": {"slug": "github-actions"}, + "check_suite": {"pull_requests": [{"number": 34}]}, + }, + "repository": {"full_name": "test/repo"}, + "sender": {"login": "github-actions[bot]", "type": "Bot"}, + } + body = json.dumps(payload).encode() + headers = {"X-GitHub-Event": "check_run", "X-GitHub-Delivery": "id-5b"} + resolver = MockResolver(github_repo_ref, github_connection) + + event = await github_adapter.parse_webhook(headers, body, resolver) + + assert event.change_request is not None + assert event.change_request.identity.native_id == 34 + + @pytest.mark.asyncio + async def test_check_suite_has_no_single_check( + self, + github_adapter: GitHubAdapter, + github_repo_ref: RepositoryRef, + github_connection: Connection, + ): + """A check_suite bundles many check runs; it doesn't map onto the + single optional CheckRun NormalizedEvent.check carries. Leave it None + -- callers needing the full set call get_checks().""" + payload = { + "action": "completed", + "check_suite": { + "id": 1, + "status": "completed", + "conclusion": "success", + "pull_requests": [{"number": 56}], + }, + "repository": {"full_name": "test/repo"}, + "sender": {"login": "github-actions[bot]", "type": "Bot"}, + } + body = json.dumps(payload).encode() + headers = {"X-GitHub-Event": "check_suite", "X-GitHub-Delivery": "id-6"} + resolver = MockResolver(github_repo_ref, github_connection) + + event = await github_adapter.parse_webhook(headers, body, resolver) + + assert event.check is None + assert event.kind == EventKind.CHECK_UPDATED + # check_suite.pull_requests[] identifies the PR even though this + # event kind never populates .check (see docstring above). + assert event.change_request is not None + assert event.change_request.identity.native_id == 56 + + @pytest.mark.asyncio + async def test_check_event_without_pull_requests_leaves_change_request_none( + self, + github_adapter: GitHubAdapter, + github_repo_ref: RepositoryRef, + github_connection: Connection, + ): + """A check event for a commit with no associated PR (e.g. a push to + main) has no pull_requests stub at all -- change_request stays None + rather than raising.""" + payload = { + "action": "completed", + "check_suite": {"id": 1, "status": "completed", "pull_requests": []}, + "repository": {"full_name": "test/repo"}, + "sender": {"login": "github-actions[bot]", "type": "Bot"}, + } + body = json.dumps(payload).encode() + headers = {"X-GitHub-Event": "check_suite", "X-GitHub-Delivery": "id-6b"} + resolver = MockResolver(github_repo_ref, github_connection) + + event = await github_adapter.parse_webhook(headers, body, resolver) + + assert event.change_request is None + + +class TestCreateBranch: + @pytest.mark.asyncio + async def test_delegates_to_client( + self, + github_adapter_with_mock_client: GitHubAdapter, + github_repo_ref: RepositoryRef, + mock_github_http_client: GitHubClient, + ): + """create_branch splits the namespace and delegates, discarding the + client's return value (the protocol returns None).""" + mock_github_http_client.create_branch = AsyncMock( + return_value={"ref": "refs/heads/forge/feature", "object": {"sha": "deadbeef"}} + ) + + result = await github_adapter_with_mock_client.create_branch( + github_repo_ref, name="forge/feature", base="main" + ) + + mock_github_http_client.create_branch.assert_awaited_once_with( + owner="test", repo="repo", branch_name="forge/feature", base="main" + ) + assert result is None diff --git a/tests/contracts/test_github_contracts.py b/tests/contracts/test_github_contracts.py deleted file mode 100644 index d42bbb6ce..000000000 --- a/tests/contracts/test_github_contracts.py +++ /dev/null @@ -1,421 +0,0 @@ -"""Contract tests for GitHub API/webhook parsing. - -These tests verify that parse_github_webhook() correctly handles -real GitHub webhook payloads. Fixtures are based on actual webhook events. -""" - -import json -from pathlib import Path - -from forge.integrations.github.webhooks import ( - is_ci_failure, - is_ci_success, - is_pr_merged, - is_pr_review_approved, - is_pr_review_changes_requested, - parse_github_webhook, -) - -FIXTURES_DIR = Path(__file__).parent / "fixtures" / "github_api_responses" - - -def load_fixture(filename: str) -> dict: - """Load a JSON fixture file.""" - with open(FIXTURES_DIR / filename) as f: - return json.load(f) - - -class TestParseCheckRunEvents: - """Test parsing of check_run webhook events.""" - - def test_parse_check_run_success(self): - """Parse a successful CI check_run event.""" - payload = load_fixture("check_run_success.json") - data = parse_github_webhook( - payload=payload, - event_type="check_run", - event_id="delivery-123", - ) - - assert data.event_id == "delivery-123" - assert data.event_type == "check_run" - assert data.action == "completed" - assert data.repo_full_name == "acme/backend" - assert data.sender_login == "github-actions[bot]" - - # Check run specific fields - assert data.check_status == "completed" - assert data.check_conclusion == "success" - assert data.commit_sha == "abc123def456789012345678901234567890" - - # PR association - assert data.pr_number == 42 - assert data.branch_name == "feature/PROJ-104-oauth" - - # Ticket extraction from branch name - assert data.ticket_key == "PROJ-104" - - # Verify helper function - assert is_ci_success(data) is True - assert is_ci_failure(data) is False - - def test_parse_check_run_failure(self): - """Parse a failed CI check_run event.""" - payload = load_fixture("check_run_failure.json") - data = parse_github_webhook( - payload=payload, - event_type="check_run", - event_id="delivery-456", - ) - - assert data.check_status == "completed" - assert data.check_conclusion == "failure" - assert data.pr_number == 45 - assert data.branch_name == "bugfix/PROJ-156-password-chars" - assert data.ticket_key == "PROJ-156" - - # Verify helper functions - assert is_ci_success(data) is False - assert is_ci_failure(data) is True - - -class TestParsePullRequestEvents: - """Test parsing of pull_request webhook events.""" - - def test_parse_pull_request_opened(self): - """Parse a pull_request opened event.""" - payload = load_fixture("pull_request_opened.json") - data = parse_github_webhook( - payload=payload, - event_type="pull_request", - event_id="delivery-789", - ) - - assert data.event_type == "pull_request" - assert data.action == "opened" - assert data.repo_full_name == "acme/backend" - - # PR specific fields - assert data.pr_number == 42 - assert data.pr_url == "https://github.com/acme/backend/pull/42" - assert data.pr_state == "open" - assert data.branch_name == "feature/PROJ-104-oauth" - - # Ticket extraction from PR title (takes precedence over branch) - assert data.ticket_key == "PROJ-104" - - def test_parse_pull_request_merged(self): - """Parse a pull_request closed+merged event.""" - payload = { - "action": "closed", - "pull_request": { - "number": 42, - "state": "closed", - "merged": True, - "title": "PROJ-104: OAuth implementation", - "head": {"ref": "feature/PROJ-104"}, - "html_url": "https://github.com/acme/backend/pull/42" - }, - "repository": {"full_name": "acme/backend"}, - "sender": {"login": "senior-dev"} - } - data = parse_github_webhook( - payload=payload, - event_type="pull_request", - event_id="delivery-merge-1", - ) - - assert data.action == "closed" - assert data.pr_state == "closed" - assert is_pr_merged(data) is True - - def test_parse_pull_request_closed_not_merged(self): - """Parse a pull_request closed without merge.""" - payload = { - "action": "closed", - "pull_request": { - "number": 43, - "state": "closed", - "merged": False, - "title": "WIP: Experimental feature", - "head": {"ref": "feature/experiment"}, - "html_url": "https://github.com/acme/backend/pull/43" - }, - "repository": {"full_name": "acme/backend"}, - "sender": {"login": "dev-user"} - } - data = parse_github_webhook( - payload=payload, - event_type="pull_request", - event_id="delivery-close-1", - ) - - assert is_pr_merged(data) is False - - -class TestParsePullRequestReviewEvents: - """Test parsing of pull_request_review webhook events.""" - - def test_parse_pr_review_approved(self): - """Parse a pull_request_review approved event.""" - payload = load_fixture("pull_request_review_approved.json") - data = parse_github_webhook( - payload=payload, - event_type="pull_request_review", - event_id="delivery-review-1", - ) - - assert data.event_type == "pull_request_review" - assert data.action == "submitted" - assert data.pr_number == 42 - assert data.pr_url == "https://github.com/acme/backend/pull/42" - assert data.branch_name == "feature/PROJ-104-oauth" - assert data.ticket_key == "PROJ-104" - assert data.sender_login == "senior-dev" - - # Verify helper function - assert is_pr_review_approved(data) is True - assert is_pr_review_changes_requested(data) is False - - def test_parse_pr_review_changes_requested(self): - """Parse a pull_request_review with changes requested.""" - payload = { - "action": "submitted", - "review": { - "id": 1876543211, - "user": {"login": "senior-dev"}, - "body": "Please add error handling for the token refresh.", - "state": "changes_requested", - "submitted_at": "2024-03-20T15:00:00Z" - }, - "pull_request": { - "number": 42, - "state": "open", - "title": "PROJ-104: OAuth implementation", - "head": {"ref": "feature/PROJ-104"}, - "html_url": "https://github.com/acme/backend/pull/42" - }, - "repository": {"full_name": "acme/backend"}, - "sender": {"login": "senior-dev"} - } - data = parse_github_webhook( - payload=payload, - event_type="pull_request_review", - event_id="delivery-review-2", - ) - - assert is_pr_review_approved(data) is False - assert is_pr_review_changes_requested(data) is True - - -class TestTicketKeyExtraction: - """Test Jira ticket key extraction from various sources.""" - - def test_extract_from_pr_title(self): - """Extract ticket from PR title.""" - payload = { - "action": "opened", - "pull_request": { - "number": 1, - "state": "open", - "title": "[PROJ-123] Fix login bug", - "head": {"ref": "fix-login"}, - "html_url": "https://github.com/org/repo/pull/1" - }, - "repository": {"full_name": "org/repo"}, - "sender": {"login": "user"} - } - data = parse_github_webhook(payload, "pull_request", "id-1") - assert data.ticket_key == "PROJ-123" - - def test_extract_from_branch_when_title_has_no_ticket(self): - """Fall back to branch name when title has no ticket.""" - payload = { - "action": "opened", - "pull_request": { - "number": 1, - "state": "open", - "title": "Fix login bug", - "head": {"ref": "feature/PROJ-456-login"}, - "html_url": "https://github.com/org/repo/pull/1" - }, - "repository": {"full_name": "org/repo"}, - "sender": {"login": "user"} - } - data = parse_github_webhook(payload, "pull_request", "id-2") - assert data.ticket_key == "PROJ-456" - - def test_extract_ticket_various_formats(self): - """Test ticket extraction with various naming conventions.""" - test_cases = [ - ("PROJ-123: Feature title", "PROJ-123"), - ("[PROJ-123] Feature title", "PROJ-123"), - ("Feature PROJ-123 implementation", "PROJ-123"), - ("feat: add PROJ-123 support", "PROJ-123"), - ("feature/PROJ-123-oauth", "PROJ-123"), - ("bugfix/proj-456-fix", "PROJ-456"), # Case insensitive - ("ABC-1", "ABC-1"), # Single digit - ("LONGPROJECT-99999", "LONGPROJECT-99999"), # Long project key - ] - - for text, expected_key in test_cases: - payload = { - "action": "opened", - "pull_request": { - "number": 1, - "state": "open", - "title": text, - "head": {"ref": "main"}, - "html_url": "https://github.com/org/repo/pull/1" - }, - "repository": {"full_name": "org/repo"}, - "sender": {"login": "user"} - } - data = parse_github_webhook(payload, "pull_request", "id") - assert data.ticket_key == expected_key, f"Failed for: {text}" - - def test_no_ticket_found(self): - """Return None when no ticket key found.""" - payload = { - "action": "opened", - "pull_request": { - "number": 1, - "state": "open", - "title": "Fix some bug", - "head": {"ref": "fix-bug"}, - "html_url": "https://github.com/org/repo/pull/1" - }, - "repository": {"full_name": "org/repo"}, - "sender": {"login": "user"} - } - data = parse_github_webhook(payload, "pull_request", "id") - assert data.ticket_key is None - - -class TestPushEvents: - """Test parsing of push webhook events.""" - - def test_parse_push_with_ticket_in_branch(self): - """Parse a push event with ticket in branch name.""" - payload = { - "ref": "refs/heads/feature/PROJ-789-feature", - "after": "newcommitsha123456789012345678901234", - "before": "oldcommitsha123456789012345678901234", - "repository": {"full_name": "acme/backend"}, - "sender": {"login": "developer"} - } - data = parse_github_webhook(payload, "push", "delivery-push-1") - - assert data.event_type == "push" - assert data.branch_name == "feature/PROJ-789-feature" - assert data.commit_sha == "newcommitsha123456789012345678901234" - assert data.ticket_key == "PROJ-789" - - -class TestEdgeCases: - """Test edge cases and missing fields.""" - - def test_minimal_payload(self): - """Handle minimal payload with missing optional fields.""" - payload = { - "action": "created", - "repository": {}, - "sender": {} - } - data = parse_github_webhook(payload, "unknown", "id-1") - - assert data.event_type == "unknown" - assert data.action == "created" - assert data.repo_full_name == "" - assert data.sender_login == "" - assert data.ticket_key is None - assert data.pr_number is None - - def test_check_run_without_pull_requests(self): - """Handle check_run without associated PRs.""" - payload = { - "action": "completed", - "check_run": { - "id": 123, - "status": "completed", - "conclusion": "success", - "head_sha": "sha123", - "pull_requests": [] # No associated PRs - }, - "repository": {"full_name": "acme/repo"}, - "sender": {"login": "bot"} - } - data = parse_github_webhook(payload, "check_run", "id-1") - - assert data.check_status == "completed" - assert data.check_conclusion == "success" - assert data.pr_number is None - assert data.ticket_key is None - - def test_raw_payload_preserved(self): - """Verify raw payload is preserved in parsed data.""" - payload = { - "action": "opened", - "custom_field": "custom_value", - "repository": {"full_name": "org/repo"}, - "sender": {"login": "user"} - } - data = parse_github_webhook(payload, "test", "id-1") - - assert data.raw_payload == payload - assert data.raw_payload["custom_field"] == "custom_value" - - -class TestIssueCommentEvents: - """Test parsing of issue_comment events (PR comments).""" - - def test_parse_pr_comment(self): - """Parse a comment on a pull request.""" - payload = { - "action": "created", - "issue": { - "number": 42, - "title": "PROJ-104: OAuth implementation", - "html_url": "https://github.com/acme/backend/pull/42", - "pull_request": { - "url": "https://api.github.com/repos/acme/backend/pulls/42" - } - }, - "comment": { - "id": 12345, - "body": "Looks good, just one question..." - }, - "repository": {"full_name": "acme/backend"}, - "sender": {"login": "reviewer"} - } - data = parse_github_webhook(payload, "issue_comment", "id-1") - - assert data.event_type == "issue_comment" - assert data.action == "created" - assert data.pr_number == 42 - assert data.pr_url == "https://github.com/acme/backend/pull/42" - assert data.ticket_key == "PROJ-104" - - def test_parse_issue_comment_not_pr(self): - """Parse a comment on a regular issue (not a PR).""" - payload = { - "action": "created", - "issue": { - "number": 100, - "title": "Bug report", - "html_url": "https://github.com/acme/backend/issues/100" - # No pull_request field - }, - "comment": { - "id": 12346, - "body": "Can you provide more details?" - }, - "repository": {"full_name": "acme/backend"}, - "sender": {"login": "maintainer"} - } - data = parse_github_webhook(payload, "issue_comment", "id-2") - - assert data.event_type == "issue_comment" - # Should not set PR fields since this is not a PR - assert data.pr_number is None - assert data.ticket_key is None diff --git a/tests/integration/api/test_webhook_queue.py b/tests/integration/api/test_webhook_queue.py index 250153280..e44a78dca 100644 --- a/tests/integration/api/test_webhook_queue.py +++ b/tests/integration/api/test_webhook_queue.py @@ -6,7 +6,7 @@ from unittest.mock import patch from forge.queue.models import QueueMessage -from forge.queue.producer import GITHUB_STREAM, JIRA_STREAM, QueueProducer +from forge.queue.producer import JIRA_STREAM, SOURCE_CONTROL_STREAM, QueueProducer from tests.fixtures.github_payloads import WEBHOOK_CHECK_RUN_COMPLETED_SUCCESS from tests.fixtures.jira_payloads import WEBHOOK_ISSUE_CREATED @@ -106,11 +106,11 @@ async def test_github_delivery_is_authenticated_queued_and_deduplicated( ) assert accepted.status_code == 202 - assert accepted.json()["status"] == "accepted" + assert accepted.json()["status"] == "queued" assert duplicate.status_code == 202 assert duplicate.json()["status"] == "duplicate" - messages = await _stream_messages(redis_client, GITHUB_STREAM) + messages = await _stream_messages(redis_client, SOURCE_CONTROL_STREAM) assert len(messages) == 1 assert messages[0].event_id == "github-delivery-1" - assert messages[0].source.value == "github" + assert messages[0].source.value == "source_control" diff --git a/tests/integration/orchestrator/test_ci_fix_attempt_status_comments.py b/tests/integration/orchestrator/test_ci_fix_attempt_status_comments.py index 2d4b5262b..827832060 100644 --- a/tests/integration/orchestrator/test_ci_fix_attempt_status_comments.py +++ b/tests/integration/orchestrator/test_ci_fix_attempt_status_comments.py @@ -69,7 +69,7 @@ async def test_first_attempt_posts_comment_with_1_of_max(self): with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.ci_evaluator.ContainerRunner", return_value=mock_runner): - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(MagicMock(), mock_github)): with patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") as mock_prepare: mock_prepare.return_value = (Path("/tmp/test-workspace"), None) with patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", AsyncMock()): @@ -116,7 +116,7 @@ async def test_second_attempt_posts_comment_with_2_of_max(self): with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.ci_evaluator.ContainerRunner", return_value=mock_runner): - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(MagicMock(), mock_github)): with patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") as mock_prepare: mock_prepare.return_value = (Path("/tmp/test-workspace"), None) with patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", AsyncMock()): @@ -163,7 +163,7 @@ async def test_final_attempt_posts_comment_with_max_of_max(self): with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.ci_evaluator.ContainerRunner", return_value=mock_runner): - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(MagicMock(), mock_github)): with patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") as mock_prepare: mock_prepare.return_value = (Path("/tmp/test-workspace"), None) with patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", AsyncMock()): @@ -210,7 +210,7 @@ async def test_comment_posted_to_feature_ticket_not_task(self): with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.ci_evaluator.ContainerRunner", return_value=mock_runner): - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(MagicMock(), mock_github)): with patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") as mock_prepare: mock_prepare.return_value = (Path("/tmp/test-workspace"), None) with patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", AsyncMock()): @@ -270,7 +270,7 @@ def capture_comment(ticket_key, message): with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.ci_evaluator.ContainerRunner", return_value=mock_runner): - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(MagicMock(), mock_github)): with patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") as mock_prepare: mock_prepare.return_value = (Path("/tmp/test-workspace"), None) with patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", AsyncMock()): @@ -317,7 +317,7 @@ async def test_different_max_attempts_values(self): with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.ci_evaluator.ContainerRunner", return_value=mock_runner): - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(MagicMock(), mock_github)): with patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") as mock_prepare: mock_prepare.return_value = (Path("/tmp/test-workspace"), None) with patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", AsyncMock()): @@ -369,7 +369,7 @@ async def test_workflow_continues_when_comment_posting_fails(self, caplog): with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.ci_evaluator.ContainerRunner", return_value=mock_runner): - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(MagicMock(), mock_github)): with patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") as mock_prepare: mock_prepare.return_value = (Path("/tmp/test-workspace"), None) with patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", AsyncMock()): @@ -415,7 +415,7 @@ async def test_jira_client_closed_even_on_comment_error(self): with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.ci_evaluator.ContainerRunner", return_value=mock_runner): - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(MagicMock(), mock_github)): with patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") as mock_prepare: mock_prepare.return_value = (Path("/tmp/test-workspace"), None) with patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", AsyncMock()): @@ -448,5 +448,5 @@ async def test_no_comment_posted_when_no_failed_checks(self): # Verify no comment posted (early return) assert mock_jira.add_comment.call_count == 0 - # Verify early return routes to human_review_gate - assert result["current_node"] == "human_review_gate" + # Defensive entry re-evaluates live CI instead of assuming success. + assert result["current_node"] == "ci_evaluator" diff --git a/tests/integration/orchestrator/test_task_handoff.py b/tests/integration/orchestrator/test_task_handoff.py index 6cceeeb62..164032233 100644 --- a/tests/integration/orchestrator/test_task_handoff.py +++ b/tests/integration/orchestrator/test_task_handoff.py @@ -54,7 +54,7 @@ async def test_workspace_setup_node_creates_forge_directory(self): patch("forge.workflow.nodes.workspace_setup.GitOperations") as MockGit, patch("forge.workflow.nodes.workspace_setup.GuardrailsLoader") as MockGuardrails, patch("forge.workflow.nodes.workspace_setup.JiraClient") as MockJira, - patch("forge.workflow.nodes.workspace_setup.GitHubClient") as MockGitHub, + patch("forge.workflow.nodes.workspace_setup.get_adapter") as MockGetAdapter, ): mock_git = MagicMock() MockGit.return_value = mock_git @@ -70,10 +70,14 @@ async def test_workspace_setup_node_creates_forge_directory(self): MockJira.return_value.get_labels = AsyncMock(return_value=[]) MockJira.return_value.add_labels = AsyncMock() MockJira.return_value.remove_labels = AsyncMock() - MockGitHub.return_value.get_repository = AsyncMock( - return_value={"default_branch": "main"} + repo_ref = MagicMock(namespace="test-org/test-repo") + adapter = MagicMock() + adapter.get_git_credentials = AsyncMock(return_value=MagicMock()) + adapter.resolve_default_branch = AsyncMock(return_value="main") + adapter.ensure_write_target = AsyncMock( + return_value=MagicMock(fork_owner="forge-bot", fork_repo="test-repo") ) - MockGitHub.return_value.close = AsyncMock() + MockGetAdapter.return_value = (repo_ref, adapter) result = await setup_workspace(initial_state) diff --git a/tests/integration/redis/test_queue_integration.py b/tests/integration/redis/test_queue_integration.py index 826aa6502..13b226d4e 100644 --- a/tests/integration/redis/test_queue_integration.py +++ b/tests/integration/redis/test_queue_integration.py @@ -13,7 +13,7 @@ from forge.models.events import EventSource from forge.queue.consumer import CONSUMER_GROUP, QueueConsumer from forge.queue.models import QueueMessage -from forge.queue.producer import GITHUB_STREAM, JIRA_STREAM, QueueProducer +from forge.queue.producer import JIRA_STREAM, SOURCE_CONTROL_STREAM, QueueProducer from forge.queue.retry import RETRY_QUEUE_KEY, RetryEntry, RetryQueue @@ -42,19 +42,19 @@ async def test_publish_jira_event(self, redis_client): assert stream_len == 1 async def test_publish_github_event(self, redis_client): - """Publish a GitHub event to the queue.""" + """Publish a source-control event to the queue.""" producer = QueueProducer(redis_client=redis_client) await producer.publish( event_id="gh-event-456", - source=EventSource.GITHUB, + source=EventSource.SOURCE_CONTROL, event_type="check_run:completed", ticket_key="TEST-456", payload={"check_run": {"conclusion": "success"}}, ) - # Verify message was published to GitHub stream - stream_len = await redis_client.xlen(GITHUB_STREAM) + # Verify message was published to source-control stream + stream_len = await redis_client.xlen(SOURCE_CONTROL_STREAM) assert stream_len == 1 async def test_publish_multiple_events(self, redis_client): @@ -356,13 +356,13 @@ async def test_empty_payload(self, redis_client): await producer.publish( event_id="empty-test", - source=EventSource.GITHUB, + source=EventSource.SOURCE_CONTROL, event_type="ping", ticket_key="", payload={}, ) - messages = await redis_client.xrange(GITHUB_STREAM, "-", "+") + messages = await redis_client.xrange(SOURCE_CONTROL_STREAM, "-", "+") message_id, data = messages[0] message = QueueMessage.from_redis(message_id, data) diff --git a/tests/integration/workflow/test_pr_ci_status_updates.py b/tests/integration/workflow/test_pr_ci_status_updates.py index 60ec43fc0..c698bde40 100644 --- a/tests/integration/workflow/test_pr_ci_status_updates.py +++ b/tests/integration/workflow/test_pr_ci_status_updates.py @@ -201,7 +201,7 @@ async def test_first_attempt_posts_comment_with_1_of_3(self): with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.ci_evaluator.ContainerRunner", return_value=mock_runner): - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(MagicMock(), mock_github)): with patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") as mock_prepare: mock_prepare.return_value = (Path("/tmp/test-workspace"), None) with patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", AsyncMock()): @@ -251,7 +251,7 @@ async def test_second_attempt_posts_comment_with_2_of_3(self): with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.ci_evaluator.ContainerRunner", return_value=mock_runner): - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(MagicMock(), mock_github)): with patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") as mock_prepare: mock_prepare.return_value = (Path("/tmp/test-workspace"), None) with patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", AsyncMock()): @@ -301,7 +301,7 @@ async def test_third_attempt_posts_comment_with_3_of_3(self): with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.ci_evaluator.ContainerRunner", return_value=mock_runner): - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(MagicMock(), mock_github)): with patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") as mock_prepare: mock_prepare.return_value = (Path("/tmp/test-workspace"), None) with patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", AsyncMock()): @@ -484,7 +484,7 @@ async def test_workflow_continues_when_ci_attempt_comment_posting_fails(self, ca with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.ci_evaluator.ContainerRunner", return_value=mock_runner): - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(MagicMock(), mock_github)): with patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") as mock_prepare: mock_prepare.return_value = (Path("/tmp/test-workspace"), None) with patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", AsyncMock()): diff --git a/tests/unit/api/routes/test_github_webhook.py b/tests/unit/api/routes/test_github_webhook.py index a9e520356..373908d11 100644 --- a/tests/unit/api/routes/test_github_webhook.py +++ b/tests/unit/api/routes/test_github_webhook.py @@ -9,6 +9,15 @@ from httpx import ASGITransport, AsyncClient from pydantic import SecretStr +from forge.config import get_settings +from forge.integrations.source_control.contracts import ( + Connection, + EventKind, + Provider, + RepositoryRef, + ResolvedRepository, +) +from forge.integrations.source_control.errors import NotFoundError, ProviderConfigError from forge.main import app from tests.fixtures.github_payloads import ( WEBHOOK_CHECK_RUN_COMPLETED_FAILURE, @@ -31,17 +40,18 @@ class TestGitHubWebhookRoute: """Tests for /api/v1/webhooks/github endpoint.""" @pytest.mark.asyncio - async def test_valid_webhook_returns_202(self): + async def test_valid_webhook_returns_202(self, mock_settings): """Valid webhook with correct signature returns 202 Accepted.""" payload = json.dumps(WEBHOOK_CHECK_RUN_COMPLETED_SUCCESS).encode() secret = "test-github-webhook-secret" signature = compute_signature(payload, secret) - mock_settings = MagicMock() - mock_settings.github_webhook_secret = SecretStr(secret) + mock_settings = mock_settings.model_copy( + update={"github_webhook_secret": SecretStr(secret)} + ) mock_producer = MagicMock() - mock_producer.publish_once = AsyncMock() + mock_producer.publish_event = AsyncMock() with ( patch("forge.api.routes.github.get_settings", return_value=mock_settings), @@ -64,12 +74,13 @@ async def test_valid_webhook_returns_202(self): assert response.status_code == 202 @pytest.mark.asyncio - async def test_invalid_signature_returns_401(self): + async def test_invalid_signature_returns_401(self, mock_settings): """Invalid signature returns 401 Unauthorized.""" payload = json.dumps(WEBHOOK_CHECK_RUN_COMPLETED_SUCCESS).encode() - mock_settings = MagicMock() - mock_settings.github_webhook_secret = SecretStr("correct-secret") + mock_settings = mock_settings.model_copy( + update={"github_webhook_secret": SecretStr("correct-secret")} + ) with patch("forge.api.routes.github.get_settings", return_value=mock_settings): async with AsyncClient( @@ -88,12 +99,13 @@ async def test_invalid_signature_returns_401(self): assert response.status_code == 401 @pytest.mark.asyncio - async def test_missing_signature_returns_401(self): + async def test_missing_signature_returns_401(self, mock_settings): """Missing signature header returns 401 when secret is configured.""" payload = json.dumps(WEBHOOK_CHECK_RUN_COMPLETED_SUCCESS).encode() - mock_settings = MagicMock() - mock_settings.github_webhook_secret = SecretStr("some-secret") + mock_settings = mock_settings.model_copy( + update={"github_webhook_secret": SecretStr("some-secret")} + ) with patch("forge.api.routes.github.get_settings", return_value=mock_settings): async with AsyncClient( @@ -111,17 +123,46 @@ async def test_missing_signature_returns_401(self): assert response.status_code == 401 @pytest.mark.asyncio - async def test_check_run_success_published(self): + async def test_missing_signature_is_accepted_when_secret_is_empty(self, mock_settings): + """Unsigned synthetic webhooks are accepted when verification is disabled.""" + payload = json.dumps(WEBHOOK_CHECK_RUN_COMPLETED_SUCCESS).encode() + mock_settings = mock_settings.model_copy(update={"github_webhook_secret": SecretStr("")}) + mock_producer = MagicMock() + mock_producer.publish_event = AsyncMock() + + with ( + patch("forge.api.routes.github.get_settings", return_value=mock_settings), + patch("forge.api.routes.github.QueueProducer", return_value=mock_producer), + ): + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://test" + ) as client: + response = await client.post( + "/api/v1/webhooks/github", + content=payload, + headers={ + "Content-Type": "application/json", + "X-GitHub-Event": "check_run", + "X-GitHub-Delivery": "unsigned-delivery", + }, + ) + + assert response.status_code == 202 + mock_producer.publish_event.assert_awaited_once() + + @pytest.mark.asyncio + async def test_check_run_success_published(self, mock_settings): """Check run success event is published.""" payload = json.dumps(WEBHOOK_CHECK_RUN_COMPLETED_SUCCESS).encode() secret = "test-github-webhook-secret" signature = compute_signature(payload, secret) - mock_settings = MagicMock() - mock_settings.github_webhook_secret = SecretStr(secret) + mock_settings = mock_settings.model_copy( + update={"github_webhook_secret": SecretStr(secret)} + ) mock_producer = MagicMock() - mock_producer.publish_once = AsyncMock() + mock_producer.publish_event = AsyncMock() with ( patch("forge.api.routes.github.get_settings", return_value=mock_settings), @@ -142,20 +183,21 @@ async def test_check_run_success_published(self): ) assert response.status_code == 202 - mock_producer.publish_once.assert_called_once() + mock_producer.publish_event.assert_called_once() @pytest.mark.asyncio - async def test_check_run_failure_published(self): + async def test_check_run_failure_published(self, mock_settings): """Check run failure event is published.""" payload = json.dumps(WEBHOOK_CHECK_RUN_COMPLETED_FAILURE).encode() secret = "test-github-webhook-secret" signature = compute_signature(payload, secret) - mock_settings = MagicMock() - mock_settings.github_webhook_secret = SecretStr(secret) + mock_settings = mock_settings.model_copy( + update={"github_webhook_secret": SecretStr(secret)} + ) mock_producer = MagicMock() - mock_producer.publish_once = AsyncMock() + mock_producer.publish_event = AsyncMock() with ( patch("forge.api.routes.github.get_settings", return_value=mock_settings), @@ -176,20 +218,21 @@ async def test_check_run_failure_published(self): ) assert response.status_code == 202 - mock_producer.publish_once.assert_called_once() + mock_producer.publish_event.assert_called_once() @pytest.mark.asyncio - async def test_pr_review_approved_published(self): + async def test_pr_review_approved_published(self, mock_settings): """PR review approved event is published.""" payload = json.dumps(WEBHOOK_PULL_REQUEST_REVIEW_APPROVED).encode() secret = "test-github-webhook-secret" signature = compute_signature(payload, secret) - mock_settings = MagicMock() - mock_settings.github_webhook_secret = SecretStr(secret) + mock_settings = mock_settings.model_copy( + update={"github_webhook_secret": SecretStr(secret)} + ) mock_producer = MagicMock() - mock_producer.publish_once = AsyncMock() + mock_producer.publish_event = AsyncMock() with ( patch("forge.api.routes.github.get_settings", return_value=mock_settings), @@ -212,7 +255,7 @@ async def test_pr_review_approved_published(self): assert response.status_code == 202 @pytest.mark.asyncio - async def test_webhook_delivery_comment_from_app_bot(self): + async def test_webhook_delivery_comment_from_app_bot(self, mock_settings): """Standard App bot comment webhook delivery is received and queued successfully.""" comment_payload = { "action": "created", @@ -236,11 +279,12 @@ async def test_webhook_delivery_comment_from_app_bot(self): secret = "test-github-webhook-secret" signature = compute_signature(payload, secret) - mock_settings = MagicMock() - mock_settings.github_webhook_secret = SecretStr(secret) + mock_settings = mock_settings.model_copy( + update={"github_webhook_secret": SecretStr(secret)} + ) mock_producer = MagicMock() - mock_producer.publish_once = AsyncMock() + mock_producer.publish_event = AsyncMock() with ( patch("forge.api.routes.github.get_settings", return_value=mock_settings), @@ -261,10 +305,10 @@ async def test_webhook_delivery_comment_from_app_bot(self): ) assert response.status_code == 202 - mock_producer.publish_once.assert_called_once() + mock_producer.publish_event.assert_called_once() @pytest.mark.asyncio - async def test_webhook_delivery_comment_from_custom_dev_pat(self): + async def test_webhook_delivery_comment_from_custom_dev_pat(self, mock_settings): """Custom dev PAT user comment webhook delivery is received and queued successfully.""" comment_payload = { "action": "created", @@ -288,11 +332,12 @@ async def test_webhook_delivery_comment_from_custom_dev_pat(self): secret = "test-github-webhook-secret" signature = compute_signature(payload, secret) - mock_settings = MagicMock() - mock_settings.github_webhook_secret = SecretStr(secret) + mock_settings = mock_settings.model_copy( + update={"github_webhook_secret": SecretStr(secret)} + ) mock_producer = MagicMock() - mock_producer.publish_once = AsyncMock() + mock_producer.publish_event = AsyncMock() with ( patch("forge.api.routes.github.get_settings", return_value=mock_settings), @@ -313,59 +358,461 @@ async def test_webhook_delivery_comment_from_custom_dev_pat(self): ) assert response.status_code == 202 - mock_producer.publish_once.assert_called_once() + mock_producer.publish_event.assert_called_once() + + +class TestWebhookRouteNormalizedEventCutover: + """Route behavior post-cutover: publishes a NormalizedEvent via + QueueProducer.publish_event, using GitHubAdapter.verify_webhook/.parse_webhook + and the process-wide Registry instead of the old parse_github_webhook path. + """ + + def _sign(self, body: bytes, secret: str) -> str: + return "sha256=" + hmac.new(secret.encode(), body, hashlib.sha256).hexdigest() + + @pytest.fixture(autouse=True) + def _reset_settings_cache(self): + # get_settings() is a process-wide @lru_cache'd singleton (forge.config), + # already primed by an earlier import of forge.main in this test session. + # These tests rely on monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", ...) to + # change what the route sees, which only works if the cache is cleared so + # the next call rebuilds Settings from the current environment. Cleared + # again on teardown so this doesn't leak a stale secret into other test + # modules that call the real get_settings(). + get_settings.cache_clear() + yield + get_settings.cache_clear() + + @pytest.mark.asyncio + async def test_pull_request_opened_publishes_normalized_event(self, monkeypatch): + payload = { + "action": "opened", + "pull_request": { + "number": 42, + "html_url": "https://github.com/acme/payments/pull/42", + "title": "PROJ-123 Add feature", + "body": "", + "state": "open", + "draft": False, + "head": {"ref": "feature/PROJ-123"}, + "base": {"ref": "main"}, + }, + "repository": {"full_name": "acme/payments"}, + "sender": {"login": "octocat", "type": "User"}, + } + body = json.dumps(payload).encode() + secret = "test-secret" + monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", secret) + + published = {} + + async def fake_publish_event(_self, event, ticket_key): + published["event"] = event + published["ticket_key"] = ticket_key + return "msg-1" + + with patch("forge.queue.producer.QueueProducer.publish_event", fake_publish_event): + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://test" + ) as client: + response = await client.post( + "/api/v1/webhooks/github", + content=body, + headers={ + "X-GitHub-Event": "pull_request", + "X-GitHub-Delivery": "delivery-1", + "X-Hub-Signature-256": self._sign(body, secret), + }, + ) + + assert response.status_code == 202 + assert published["ticket_key"] == "PROJ-123" + assert published["event"].kind == EventKind.CR_OPENED + assert published["event"].repo_ref.namespace == "acme/payments" + + @pytest.mark.asyncio + async def test_push_event_extracts_ticket_key_from_ref(self, monkeypatch): + """Push events carry no change_request, so the ticket key must be + recovered from the pushed branch ref instead of being dropped.""" + payload = { + "ref": "refs/heads/forge/PROJ-123-add-feature", + "repository": {"full_name": "acme/payments"}, + "sender": {"login": "octocat", "type": "User"}, + } + body = json.dumps(payload).encode() + secret = "test-secret" + monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", secret) + + published = {} + + async def fake_publish_event(_self, event, ticket_key): + published["event"] = event + published["ticket_key"] = ticket_key + return "msg-push-1" + + with patch("forge.queue.producer.QueueProducer.publish_event", fake_publish_event): + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://test" + ) as client: + response = await client.post( + "/api/v1/webhooks/github", + content=body, + headers={ + "X-GitHub-Event": "push", + "X-GitHub-Delivery": "delivery-push-1", + "X-Hub-Signature-256": self._sign(body, secret), + }, + ) + + assert response.status_code == 202 + assert published["ticket_key"] == "PROJ-123" + assert published["event"].kind == EventKind.PUSH + @pytest.mark.asyncio + async def test_check_suite_without_pr_stub_extracts_ticket_key_from_head_branch( + self, monkeypatch + ): + """check_suite can fire before GitHub attaches a pull_requests stub; + the ticket key must still be recovered from head_branch so CI that + starts before PR creation isn't dropped.""" + payload = { + "action": "completed", + "check_suite": { + "head_branch": "forge/PROJ-456-add-feature", + "head_sha": "abc123", + "pull_requests": [], + }, + "repository": {"full_name": "acme/payments"}, + "sender": {"login": "octocat", "type": "Bot"}, + } + body = json.dumps(payload).encode() + secret = "test-secret" + monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", secret) -class TestGitHubWebhookParsing: - """Tests for GitHub webhook payload parsing via parse_github_webhook.""" + published = {} - def test_extract_pr_number(self): - """Extract PR number from check_run webhook.""" - from forge.integrations.github.webhooks import parse_github_webhook + async def fake_publish_event(_self, event, ticket_key): + published["event"] = event + published["ticket_key"] = ticket_key + return "msg-check-suite-1" + + with patch("forge.queue.producer.QueueProducer.publish_event", fake_publish_event): + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://test" + ) as client: + response = await client.post( + "/api/v1/webhooks/github", + content=body, + headers={ + "X-GitHub-Event": "check_suite", + "X-GitHub-Delivery": "delivery-check-suite-1", + "X-Hub-Signature-256": self._sign(body, secret), + }, + ) - data = parse_github_webhook(WEBHOOK_CHECK_RUN_COMPLETED_SUCCESS, "check_run", "evt-001") - assert data.pr_number == 42 + assert response.status_code == 202 + assert published["ticket_key"] == "PROJ-456" + assert published["event"].kind == EventKind.CHECK_UPDATED - def test_extract_check_conclusion(self): - """Extract check run conclusion.""" - from forge.integrations.github.webhooks import parse_github_webhook + @pytest.mark.asyncio + async def test_invalid_signature_returns_401(self, monkeypatch): + monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", "test-secret") + body = json.dumps({"action": "opened"}).encode() + + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/v1/webhooks/github", + content=body, + headers={ + "X-GitHub-Event": "pull_request", + "X-GitHub-Delivery": "delivery-2", + "X-Hub-Signature-256": "sha256=invalid", + }, + ) - success_data = parse_github_webhook( - WEBHOOK_CHECK_RUN_COMPLETED_SUCCESS, "check_run", "evt-001" - ) - failure_data = parse_github_webhook( - WEBHOOK_CHECK_RUN_COMPLETED_FAILURE, "check_run", "evt-002" + assert response.status_code == 401 + + @pytest.mark.asyncio + async def test_per_connection_webhook_secret_used_for_verification(self, monkeypatch): + """A repo resolved to a connection with its own webhook_secret_env is + verified against that connection's secret, not the global default -- + and a signature computed with the global secret must be rejected.""" + payload = { + "action": "opened", + "pull_request": { + "number": 1, + "html_url": "x", + "title": "t", + "body": "", + "state": "open", + "draft": False, + "head": {"ref": "a"}, + "base": {"ref": "main"}, + }, + "repository": {"full_name": "acme/custom-repo"}, + "sender": {"login": "x", "type": "User"}, + } + body = json.dumps(payload).encode() + global_secret = "global-secret" + custom_secret = "custom-connection-secret" + monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", global_secret) + monkeypatch.setenv("CUSTOM_ORG_WEBHOOK_SECRET", custom_secret) + + resolved = ResolvedRepository( + repo_ref=RepositoryRef( + id="custom-repo", + provider=Provider.GITHUB, + connection="custom-org", + namespace="acme/custom-repo", + default_branch="main", + change_request_mode="fork", + ), + connection=Connection( + name="custom-org", + provider=Provider.GITHUB, + base_url="https://api.github.com", + credential_env="GITHUB_TOKEN", + webhook_secret_env="CUSTOM_ORG_WEBHOOK_SECRET", + ), ) - assert success_data.check_conclusion == "success" - assert failure_data.check_conclusion == "failure" + async def fake_publish_event(_self, _event, _ticket_key): + return "msg-1" + + with ( + patch( + "forge.integrations.source_control.registry.Registry.resolve", + return_value=resolved, + ), + patch("forge.queue.producer.QueueProducer.publish_event", fake_publish_event), + ): + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://test" + ) as client: + accepted = await client.post( + "/api/v1/webhooks/github", + content=body, + headers={ + "X-GitHub-Event": "pull_request", + "X-GitHub-Delivery": "delivery-custom-1", + "X-Hub-Signature-256": self._sign(body, custom_secret), + }, + ) + rejected = await client.post( + "/api/v1/webhooks/github", + content=body, + headers={ + "X-GitHub-Event": "pull_request", + "X-GitHub-Delivery": "delivery-custom-2", + "X-Hub-Signature-256": self._sign(body, global_secret), + }, + ) + + assert accepted.status_code == 202 + assert rejected.status_code == 401 + + @pytest.mark.asyncio + async def test_unmanaged_repository_acks_and_drops(self, monkeypatch): + """resolve() raising NotFoundError -- the repo isn't managed by Forge -- + must ack (202) and drop, not error.""" + payload = { + "action": "opened", + "pull_request": { + "number": 1, + "html_url": "x", + "title": "t", + "body": "", + "state": "open", + "draft": False, + "head": {"ref": "a"}, + "base": {"ref": "main"}, + }, + "repository": {"full_name": "unmanaged/repo"}, + "sender": {"login": "x", "type": "User"}, + } + body = json.dumps(payload).encode() + secret = "test-secret" + monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", secret) + + with patch( + "forge.integrations.source_control.registry.Registry.resolve", + side_effect=NotFoundError("unmanaged"), + ): + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://test" + ) as client: + response = await client.post( + "/api/v1/webhooks/github", + content=body, + headers={ + "X-GitHub-Event": "pull_request", + "X-GitHub-Delivery": "delivery-3", + "X-Hub-Signature-256": self._sign(body, secret), + }, + ) + + assert response.status_code == 202 + assert response.json()["status"] == "ignored" + + @pytest.mark.asyncio + async def test_misconfigured_connection_acks_and_drops(self, monkeypatch): + """adapter.parse_webhook's internal resolver.resolve() call can raise + ProviderConfigError (distinct from NotFoundError) when a repo resolves + to a connection that isn't usable (e.g. no credential configured). + This must also ack (202) and drop, not fall through to a 500 that + GitHub would treat as retryable.""" + payload = { + "action": "opened", + "pull_request": { + "number": 1, + "html_url": "x", + "title": "t", + "body": "", + "state": "open", + "draft": False, + "head": {"ref": "a"}, + "base": {"ref": "main"}, + }, + "repository": {"full_name": "misconfigured/repo"}, + "sender": {"login": "x", "type": "User"}, + } + body = json.dumps(payload).encode() + secret = "test-secret" + monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", secret) + + with patch( + "forge.integrations.source_control.registry.Registry.resolve", + side_effect=ProviderConfigError("credential not set"), + ): + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://test" + ) as client: + response = await client.post( + "/api/v1/webhooks/github", + content=body, + headers={ + "X-GitHub-Event": "pull_request", + "X-GitHub-Delivery": "delivery-3b", + "X-Hub-Signature-256": self._sign(body, secret), + }, + ) + + assert response.status_code == 202 + assert response.json()["status"] == "ignored" + + @pytest.mark.asyncio + async def test_malformed_json_returns_400(self, monkeypatch): + """A malformed body must surface as 400, not the generic 500 handler. + + adapter.parse_webhook raises json.JSONDecodeError internally when given + a malformed body, but the route decodes the body itself first + specifically to preserve the pre-cutover 400-on-malformed-JSON behavior. + """ + secret = "test-secret" + monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", secret) + body = b"{not valid json" + + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/v1/webhooks/github", + content=body, + headers={ + "X-GitHub-Event": "pull_request", + "X-GitHub-Delivery": "delivery-4", + "X-Hub-Signature-256": self._sign(body, secret), + }, + ) + + assert response.status_code == 400 + assert response.json()["detail"] == "Invalid JSON payload" + + @pytest.mark.asyncio + async def test_duplicate_event_returns_duplicate_status(self, monkeypatch): + """publish_event returning None (a duplicate delivery id already seen) + must surface as a 202 "duplicate" ack, not silently look like success.""" + payload = { + "action": "opened", + "pull_request": { + "number": 1, + "html_url": "x", + "title": "t", + "body": "", + "state": "open", + "draft": False, + "head": {"ref": "a"}, + "base": {"ref": "main"}, + }, + "repository": {"full_name": "acme/payments"}, + "sender": {"login": "x", "type": "User"}, + } + body = json.dumps(payload).encode() + secret = "test-secret" + monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", secret) - def test_extract_repository(self): - """Extract repository from webhook.""" - from forge.integrations.github.webhooks import parse_github_webhook + with patch( + "forge.queue.producer.QueueProducer.publish_event", + AsyncMock(return_value=None), + ): + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://test" + ) as client: + response = await client.post( + "/api/v1/webhooks/github", + content=body, + headers={ + "X-GitHub-Event": "pull_request", + "X-GitHub-Delivery": "delivery-5", + "X-Hub-Signature-256": self._sign(body, secret), + }, + ) - data = parse_github_webhook(WEBHOOK_CHECK_RUN_COMPLETED_SUCCESS, "check_run", "evt-001") - assert data.repo_full_name == "org/repo" + assert response.status_code == 202 + assert response.json()["status"] == "duplicate" - def test_extract_review_state(self): - """Extract review state from PR review webhook.""" - # Review state is in the raw payload's review object - review = WEBHOOK_PULL_REQUEST_REVIEW_APPROVED.get("review", {}) - assert review.get("state") == "approved" + @pytest.mark.asyncio + async def test_missing_delivery_header_fallback_flows_into_published_event(self, monkeypatch): + """When X-GitHub-Delivery is absent, the route's generated fallback id + must be the id actually published, not just echoed in the response -- + otherwise the queued event and what's logged/returned diverge.""" + payload = { + "action": "opened", + "pull_request": { + "number": 1, + "html_url": "x", + "title": "t", + "body": "", + "state": "open", + "draft": False, + "head": {"ref": "a"}, + "base": {"ref": "main"}, + }, + "repository": {"full_name": "acme/payments"}, + "sender": {"login": "x", "type": "User"}, + } + body = json.dumps(payload).encode() + secret = "test-secret" + monkeypatch.setenv("GITHUB_WEBHOOK_SECRET", secret) - def test_extract_ci_output(self): - """Extract CI output from failed check run (via payload).""" - check_run = WEBHOOK_CHECK_RUN_COMPLETED_FAILURE.get("check_run", {}) - output = check_run.get("output", {}) - text = output.get("text", "") + published = {} - assert "test_login_validation" in text - assert "AssertionError" in text + async def fake_publish_event(_self, event, _ticket_key): + published["event"] = event + return "msg-1" - def test_detect_forge_branch(self): - """Detect ticket key from branch name.""" - from forge.integrations.github.webhooks import parse_github_webhook + with patch("forge.queue.producer.QueueProducer.publish_event", fake_publish_event): + async with AsyncClient( + transport=ASGITransport(app=app), base_url="http://test" + ) as client: + response = await client.post( + "/api/v1/webhooks/github", + content=body, + headers={ + "X-GitHub-Event": "pull_request", + "X-Hub-Signature-256": self._sign(body, secret), + }, + ) - data = parse_github_webhook(WEBHOOK_CHECK_RUN_COMPLETED_SUCCESS, "check_run", "evt-001") - # Branch is feature/TEST-123, should extract ticket key - assert data.ticket_key == "TEST-123" + assert response.status_code == 202 + response_event_id = response.json()["event_id"] + assert response_event_id != "" + assert published["event"].id == response_event_id diff --git a/tests/unit/integrations/github/test_review_threads.py b/tests/unit/integrations/github/test_review_threads.py index 5e7baa2a8..f73fa011c 100644 --- a/tests/unit/integrations/github/test_review_threads.py +++ b/tests/unit/integrations/github/test_review_threads.py @@ -1,4 +1,4 @@ -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, call import pytest @@ -120,3 +120,50 @@ async def test_review_threads_fall_back_to_rest() -> None: assert threads[0]["thread_id"] == "rest-77" assert threads[0]["comments"][0]["comment_id"] == 77 http.get.assert_awaited_once_with("/repos/org/repo/pulls/9/comments", params={"per_page": 100}) + + +@pytest.mark.asyncio +async def test_get_reviews_returns_review_submissions() -> None: + payload = [ + {"id": 1, "state": "APPROVED", "body": "LGTM", "user": {"login": "reviewer"}}, + {"id": 2, "state": "CHANGES_REQUESTED", "body": "", "user": {"login": "other"}}, + ] + response = AsyncMock() + response.raise_for_status = lambda: None + response.json = lambda: payload + http = AsyncMock() + http.get.return_value = response + github = GitHubClient() + github._get_client = AsyncMock(return_value=http) + + reviews = await github.get_reviews("org", "repo", 9) + + assert reviews == payload + http.get.assert_awaited_once_with( + "/repos/org/repo/pulls/9/reviews", params={"per_page": 100, "page": 1} + ) + + +@pytest.mark.asyncio +async def test_get_reviews_fetches_all_pages() -> None: + first_page = [{"id": review_id} for review_id in range(100)] + second_page = [{"id": 100}] + + first_response = AsyncMock() + first_response.raise_for_status = lambda: None + first_response.json = lambda: first_page + second_response = AsyncMock() + second_response.raise_for_status = lambda: None + second_response.json = lambda: second_page + http = AsyncMock() + http.get.side_effect = [first_response, second_response] + github = GitHubClient() + github._get_client = AsyncMock(return_value=http) + + reviews = await github.get_reviews("org", "repo", 9) + + assert reviews == first_page + second_page + assert http.get.await_args_list == [ + call("/repos/org/repo/pulls/9/reviews", params={"per_page": 100, "page": 1}), + call("/repos/org/repo/pulls/9/reviews", params={"per_page": 100, "page": 2}), + ] diff --git a/tests/unit/integrations/source_control/__init__.py b/tests/unit/integrations/source_control/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/unit/integrations/source_control/github/test_adapter_ci_artifacts.py b/tests/unit/integrations/source_control/github/test_adapter_ci_artifacts.py new file mode 100644 index 000000000..c198d1b47 --- /dev/null +++ b/tests/unit/integrations/source_control/github/test_adapter_ci_artifacts.py @@ -0,0 +1,52 @@ +from unittest.mock import AsyncMock + +import pytest + +from forge.integrations.source_control.contracts import ( + CheckConclusion, CheckRun, CheckStatus, Connection, Provider, RepositoryRef, +) +from forge.integrations.source_control.github.adapter import GitHubAdapter + + +def _conn(): + return Connection(name="c", provider=Provider.GITHUB, base_url="", + credential_env="GITHUB_TOKEN", webhook_secret_env="") + + +def _repo_ref(): + return RepositoryRef(id="acme/widgets", provider=Provider.GITHUB, connection="c", + namespace="acme/widgets", default_branch="main", + change_request_mode="fork") + + +def test_map_check_run_carries_output(): + adapter = GitHubAdapter(connection=_conn(), client=AsyncMock()) + entry = {"name": "pytest", "status": "completed", "conclusion": "failure", + "html_url": "h", "app": {"slug": "github-actions"}, + "details_url": "https://github.com/o/r/actions/runs/55", + "output": {"title": "T", "summary": "S", "text": "boom"}} + check = adapter._map_check_run(entry) + assert check.output == {"title": "T", "summary": "S", "text": "boom"} + + +@pytest.mark.asyncio +async def test_get_check_artifacts_returns_named_zips(): + client = AsyncMock() + client.get_run_artifacts.return_value = [{"id": 1, "name": "logs"}] + client.download_artifact_zip.return_value = b"PK\x03\x04zip" + adapter = GitHubAdapter(connection=_conn(), client=client) + check = CheckRun(name="pytest", status=CheckStatus.COMPLETED, + conclusion=CheckConclusion.FAILURE, logs_url="55") + + artifacts = await adapter.get_check_artifacts(_repo_ref(), check) + + assert artifacts == [("logs", b"PK\x03\x04zip")] + client.get_run_artifacts.assert_awaited_once_with("acme", "widgets", 55) + + +@pytest.mark.asyncio +async def test_get_check_artifacts_empty_when_no_run(): + adapter = GitHubAdapter(connection=_conn(), client=AsyncMock()) + check = CheckRun(name="prow", status=CheckStatus.COMPLETED, + conclusion=CheckConclusion.FAILURE, logs_url=None) + assert await adapter.get_check_artifacts(_repo_ref(), check) == [] diff --git a/tests/unit/integrations/source_control/github/test_adapter_fork_identity.py b/tests/unit/integrations/source_control/github/test_adapter_fork_identity.py new file mode 100644 index 000000000..95c7b42cd --- /dev/null +++ b/tests/unit/integrations/source_control/github/test_adapter_fork_identity.py @@ -0,0 +1,127 @@ +from unittest.mock import AsyncMock + +import pytest + +from forge.integrations.source_control.contracts import ( + Connection, + Provider, + RepositoryRef, + WriteTarget, +) +from forge.integrations.source_control.github.adapter import GitHubAdapter + + +def _repo_ref(mode="fork"): + return RepositoryRef( + id="acme/widgets", + provider=Provider.GITHUB, + connection="c", + namespace="acme/widgets", + default_branch="main", + change_request_mode=mode, + ) + + +def _conn(): + return Connection( + name="c", + provider=Provider.GITHUB, + base_url="", + credential_env="GITHUB_TOKEN", + webhook_secret_env="", + ) + + +@pytest.mark.asyncio +async def test_ensure_write_target_exposes_fork_identity(): + client = AsyncMock() + client.get_or_create_fork.return_value = { + "owner": {"login": "forkuser"}, + "name": "widgets", + "clone_url": "https://github.com/forkuser/widgets.git", + } + client.sync_fork_with_upstream.return_value = True + adapter = GitHubAdapter(connection=_conn(), client=client) + + target = await adapter.ensure_write_target(_repo_ref()) + + assert target.fork_owner == "forkuser" + assert target.fork_repo == "widgets" + + +@pytest.mark.asyncio +async def test_direct_mode_has_no_fork_identity(): + adapter = GitHubAdapter(connection=_conn(), client=AsyncMock()) + target = await adapter.ensure_write_target(_repo_ref(mode="direct")) + assert target.fork_owner is None and target.fork_repo is None + + +@pytest.mark.asyncio +async def test_create_change_request_uses_forkowner_colon_branch_head(): + client = AsyncMock() + client.create_pull_request.return_value = type( + "R", + (), + { + "pr": { + "number": 7, + "html_url": "u", + "title": "t", + "body": "", + "state": "open", + "head": {"ref": "feature/x"}, + "base": {"ref": "main"}, + "draft": False, + }, + "created": True, + }, + )() + adapter = GitHubAdapter(connection=_conn(), client=client) + target = WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="feature/x", + base_branch="main", + fork_owner="forkuser", + fork_repo="widgets", + ) + + await adapter.create_change_request(_repo_ref(), target, title="t", body="b") + + _, kwargs = client.create_pull_request.call_args + assert kwargs["head"] == "forkuser:feature/x" + + +@pytest.mark.asyncio +async def test_create_change_request_propagates_created_false_for_existing_pr(): + client = AsyncMock() + client.create_pull_request.return_value = type( + "R", + (), + { + "pr": { + "number": 7, + "html_url": "u", + "title": "t", + "body": "", + "state": "open", + "head": {"ref": "feature/x"}, + "base": {"ref": "main"}, + "draft": False, + }, + "created": False, + }, + )() + adapter = GitHubAdapter(connection=_conn(), client=client) + target = WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="feature/x", + base_branch="main", + fork_owner="forkuser", + fork_repo="widgets", + ) + + change_request = await adapter.create_change_request(_repo_ref(), target, title="t", body="b") + + assert change_request.created is False diff --git a/tests/unit/integrations/source_control/github/test_adapter_review_threads.py b/tests/unit/integrations/source_control/github/test_adapter_review_threads.py new file mode 100644 index 000000000..a14d5b076 --- /dev/null +++ b/tests/unit/integrations/source_control/github/test_adapter_review_threads.py @@ -0,0 +1,88 @@ +from unittest.mock import AsyncMock + +import pytest + +from forge.integrations.source_control.contracts import ( + ChangeRequestIdentity, + Connection, + Provider, + RepositoryRef, +) +from forge.integrations.source_control.github.adapter import GitHubAdapter + + +def _conn(): + return Connection( + name="c", + provider=Provider.GITHUB, + base_url="", + credential_env="GITHUB_TOKEN", + webhook_secret_env="", + ) + + +def _repo_ref(): + return RepositoryRef( + id="acme/widgets", + provider=Provider.GITHUB, + connection="c", + namespace="acme/widgets", + default_branch="main", + change_request_mode="fork", + ) + + +@pytest.mark.asyncio +async def test_get_review_thread_comments_maps_threads(): + client = AsyncMock() + client.get_pull_request_review_threads.return_value = [ + { + "thread_id": "T1", + "path": "a.py", + "line": 12, + "comments": [{"comment_id": 99, "body": "fix this", "author": "rev"}], + }, + ] + adapter = GitHubAdapter(connection=_conn(), client=client) + identity = ChangeRequestIdentity(connection="c", repository_id="acme/widgets", native_id=7) + + reviews = await adapter.get_review_thread_comments(_repo_ref(), identity) + + assert len(reviews) == 1 + review = reviews[0] + assert review.id == "T1" + assert len(review.comments) == 1 + c = review.comments[0] + assert (c.id, c.body, c.author, c.path, c.line) == ("99", "fix this", "rev", "a.py", 12) + client.get_pull_request_review_threads.assert_awaited_once_with("acme", "widgets", 7) + + +@pytest.mark.asyncio +async def test_get_review_comments_for_submission_falls_back_to_original_line(): + client = AsyncMock() + client.get_review_comments.return_value = [ + { + "id": 101, + "body": "outdated diff comment", + "user": {"login": "rev"}, + "path": "a.py", + "line": None, + "original_line": 42, + "position": None, + }, + { + "id": 102, + "body": "current diff comment", + "user": {"login": "rev"}, + "path": "b.py", + "line": 7, + "original_line": 3, + }, + ] + adapter = GitHubAdapter(connection=_conn(), client=client) + identity = ChangeRequestIdentity(connection="c", repository_id="acme/widgets", native_id=7) + + comments = await adapter.get_review_comments_for_submission(_repo_ref(), identity, "55") + + assert [(c.id, c.line) for c in comments] == [("101", 42), ("102", 7)] + client.get_review_comments.assert_awaited_once_with("acme", "widgets", 7, 55) diff --git a/tests/unit/integrations/source_control/github/test_factory_registration.py b/tests/unit/integrations/source_control/github/test_factory_registration.py new file mode 100644 index 000000000..a3385e6f2 --- /dev/null +++ b/tests/unit/integrations/source_control/github/test_factory_registration.py @@ -0,0 +1,25 @@ +from forge.integrations.source_control.contracts import Connection, Provider +from forge.integrations.source_control.github import GitHubAdapter +from forge.integrations.source_control.registry import _ADAPTER_FACTORIES + + +def test_github_factory_registered_on_import(): + factory = _ADAPTER_FACTORIES.get(Provider.GITHUB) + assert factory is not None + conn = Connection( + name="c", + provider=Provider.GITHUB, + base_url="https://api.github.com", + credential_env="GITHUB_TOKEN", + webhook_secret_env="GITHUB_WEBHOOK_SECRET", + ) + adapter = factory(conn) + assert isinstance(adapter, GitHubAdapter) + + +def test_registry_resolve_returns_github_adapter(mock_settings): + from forge.integrations.source_control.registry import load_registry + + registry = load_registry(config_path="/nonexistent/repos.yaml", settings=mock_settings) + resolved = registry.resolve("acme/widgets") + assert isinstance(resolved.adapter, GitHubAdapter) diff --git a/tests/unit/integrations/source_control/test_contracts.py b/tests/unit/integrations/source_control/test_contracts.py new file mode 100644 index 000000000..aa2f18600 --- /dev/null +++ b/tests/unit/integrations/source_control/test_contracts.py @@ -0,0 +1,91 @@ +"""Tests for source control contract data models.""" + +from datetime import UTC, datetime + +from forge.integrations.source_control.contracts import ( + Actor, + ChangeRequestIdentity, + CheckConclusion, + CheckRun, + CheckStatus, + EventKind, + NormalizedEvent, + Provider, + RepositoryRef, + Review, + ReviewComment, + ReviewState, +) + + +def _repo_ref() -> RepositoryRef: + return RepositoryRef( + id="payments-api", + provider=Provider.GITHUB, + connection="public-github", + namespace="acme/payments", + default_branch="main", + change_request_mode="fork", + ) + + +def test_repository_ref_holds_identity_and_mode(): + ref = _repo_ref() + assert ref.id == "payments-api" + assert ref.provider is Provider.GITHUB + assert ref.change_request_mode == "fork" + + +def test_change_request_identity_equality_is_by_value(): + a = ChangeRequestIdentity( + connection="public-github", repository_id="payments-api", native_id=42 + ) + b = ChangeRequestIdentity( + connection="public-github", repository_id="payments-api", native_id=42 + ) + c = ChangeRequestIdentity( + connection="public-github", repository_id="payments-api", native_id=None + ) + assert a == b + assert a != c + + +def test_change_request_identity_native_id_defaults_to_none(): + identity = ChangeRequestIdentity(connection="public-github", repository_id="payments-api") + assert identity.native_id is None + + +def test_normalized_event_only_populates_the_relevant_payload(): + event = NormalizedEvent( + id="evt-1", + kind=EventKind.CHECK_UPDATED, + repo_ref=_repo_ref(), + actor=Actor(login="forge-bot", is_bot=True), + received_at=datetime(2026, 8, 18, tzinfo=UTC), + check=CheckRun( + name="CI / Tests", status=CheckStatus.COMPLETED, conclusion=CheckConclusion.SUCCESS + ), + ) + assert event.kind is EventKind.CHECK_UPDATED + assert event.check is not None + assert event.check.conclusion is CheckConclusion.SUCCESS + assert event.change_request is None + assert event.review is None + assert event.comment is None + + +def test_review_carries_its_comments(): + review = Review( + id="rev-1", + state=ReviewState.CHANGES_REQUESTED, + body="Please fix the null check", + author="reviewer1", + comments=[ + ReviewComment( + id="c1", body="null check missing", author="reviewer1", path="a.py", line=10 + ) + ], + ) + assert review.state is ReviewState.CHANGES_REQUESTED + assert len(review.comments) == 1 + assert review.comments[0].path == "a.py" diff --git a/tests/unit/integrations/source_control/test_errors.py b/tests/unit/integrations/source_control/test_errors.py new file mode 100644 index 000000000..6049e88ff --- /dev/null +++ b/tests/unit/integrations/source_control/test_errors.py @@ -0,0 +1,39 @@ +"""Tests for the source control exception hierarchy.""" + +import pytest + +from forge.integrations.source_control.errors import ( + AuthenticationError, + ConflictError, + NotFoundError, + ProviderConfigError, + RateLimitedError, + SourceControlError, + TransientProviderError, +) + + +@pytest.mark.parametrize( + "exc_class", + [ + AuthenticationError, + ConflictError, + NotFoundError, + ProviderConfigError, + RateLimitedError, + TransientProviderError, + ], +) +def test_every_concrete_error_is_a_source_control_error(exc_class): + assert issubclass(exc_class, SourceControlError) + + +def test_rate_limited_error_carries_retry_after(): + error = RateLimitedError("rate limited", retry_after=30.0) + assert error.retry_after == 30.0 + assert str(error) == "rate limited" + + +def test_rate_limited_error_retry_after_defaults_to_none(): + error = RateLimitedError("rate limited") + assert error.retry_after is None diff --git a/tests/unit/integrations/source_control/test_protocol.py b/tests/unit/integrations/source_control/test_protocol.py new file mode 100644 index 000000000..96015f6a9 --- /dev/null +++ b/tests/unit/integrations/source_control/test_protocol.py @@ -0,0 +1,129 @@ +"""Tests for the SourceControlProvider protocol boundary.""" + +from forge.integrations.source_control.contracts import ( + Connection, + Provider, + RepositoryRef, + ResolvedRepository, + SourceControlProvider, +) + + +class _CompleteFakeProvider: + """Implements every SourceControlProvider method (bodies are irrelevant to this test).""" + + async def verify_webhook(self, _headers: object, _body: object) -> bool: + return True + + async def parse_webhook(self, _headers: object, _body: object, _resolver: object) -> object: + raise NotImplementedError + + async def resolve_default_branch(self, _repo_ref: object) -> object: + raise NotImplementedError + + async def get_git_credentials(self, _repo_ref: object) -> object: + raise NotImplementedError + + async def ensure_write_target(self, _repo_ref: object) -> object: + raise NotImplementedError + + async def create_change_request( + self, _repo_ref: object, _target: object, _title: object, _body: object, draft: bool = False + ) -> object: + raise NotImplementedError + + async def get_change_request(self, _repo_ref: object, _identity: object) -> object: + raise NotImplementedError + + async def update_change_request( + self, + _repo_ref: object, + _identity: object, + *, + title: object = None, + body: object = None, + state: object = None, + ) -> object: + raise NotImplementedError + + async def create_comment(self, _repo_ref: object, _identity: object, _body: object) -> object: + raise NotImplementedError + + async def reply_to_comment( + self, _repo_ref: object, _identity: object, _comment_id: object, _body: object + ) -> object: + raise NotImplementedError + + async def get_review_threads(self, _repo_ref: object, _identity: object) -> object: + raise NotImplementedError + + async def get_review_thread_comments(self, _repo_ref: object, _identity: object) -> object: + raise NotImplementedError + + async def get_review_comments_for_submission( + self, _repo_ref: object, _identity: object, _review_id: object + ) -> object: + raise NotImplementedError + + async def get_checks(self, _repo_ref: object, _ref: object) -> object: + raise NotImplementedError + + async def get_check_logs(self, _repo_ref: object, _check: object) -> object: + raise NotImplementedError + + async def get_check_artifacts(self, _repo_ref: object, _check: object) -> object: + raise NotImplementedError + + async def get_file(self, _repo_ref: object, _path: object, _ref: object) -> object: + raise NotImplementedError + + async def put_file( + self, _repo_ref: object, _path: object, _content: object, _message: object, _branch: object + ) -> None: + raise NotImplementedError + + async def create_branch(self, _repo_ref: object, _name: object, _base: object) -> None: + raise NotImplementedError + + async def get_authenticated_identity(self, _repo_ref: object) -> object: + raise NotImplementedError + + async def close(self) -> None: + pass + + +class _IncompleteFakeProvider: + """Missing every method except verify_webhook — must NOT satisfy the protocol.""" + + async def verify_webhook(self, _headers: object, _body: object) -> bool: + return True + + +def test_complete_fake_satisfies_the_protocol(): + assert isinstance(_CompleteFakeProvider(), SourceControlProvider) + + +def test_incomplete_fake_does_not_satisfy_the_protocol(): + assert not isinstance(_IncompleteFakeProvider(), SourceControlProvider) + + +def test_resolved_repository_defaults_to_no_adapter(): + repo_ref = RepositoryRef( + id="acme/payments", + provider=Provider.GITHUB, + connection="github-default", + namespace="acme/payments", + default_branch="main", + change_request_mode="fork", + ) + connection = Connection( + name="github-default", + provider=Provider.GITHUB, + base_url="https://api.github.com", + credential_env="GITHUB_TOKEN", + webhook_secret_env="GITHUB_WEBHOOK_SECRET", + ) + + resolved = ResolvedRepository(repo_ref=repo_ref, connection=connection) + + assert resolved.adapter is None diff --git a/tests/unit/integrations/source_control/test_registry_config.py b/tests/unit/integrations/source_control/test_registry_config.py new file mode 100644 index 000000000..d480ccd5a --- /dev/null +++ b/tests/unit/integrations/source_control/test_registry_config.py @@ -0,0 +1,271 @@ +"""Tests for repos.yaml config loading and validation.""" + +import pytest + +from forge.integrations.source_control.contracts import Provider +from forge.integrations.source_control.errors import ProviderConfigError +from forge.integrations.source_control.registry import load_registry + + +def _write_config(tmp_path, content: str): + path = tmp_path / "repos.yaml" + path.write_text(content) + return path + + +def test_missing_config_file_loads_an_empty_registry(tmp_path, mock_settings): + registry = load_registry(config_path=tmp_path / "does-not-exist.yaml", settings=mock_settings) + + assert registry.get_repository("anything") is None + assert registry.get_connection("anything") is None + + +def test_loads_a_connection_and_repository(tmp_path, mock_settings, monkeypatch): + monkeypatch.setenv("ACME_GITLAB_TOKEN", "secret") + config = _write_config( + tmp_path, + """ +connections: + acme-gitlab: + provider: gitlab + base_url: https://gitlab.acme.example.com + credential_env: ACME_GITLAB_TOKEN + webhook_secret_env: ACME_GITLAB_WEBHOOK_SECRET + allowed_namespaces: ["platform/payments"] + +repositories: + payments-api: + provider: gitlab + connection: acme-gitlab + namespace: platform/payments + default_branch: main + change_request_mode: direct +""", + ) + + registry = load_registry(config_path=config, settings=mock_settings) + + connection = registry.get_connection("acme-gitlab") + assert connection is not None + assert connection.provider is Provider.GITLAB + assert connection.credential_env == "ACME_GITLAB_TOKEN" + + repo = registry.get_repository("payments-api") + assert repo is not None + assert repo.namespace == "platform/payments" + assert repo.change_request_mode == "direct" + + +def test_rejects_unknown_provider(tmp_path, mock_settings): + config = _write_config( + tmp_path, + """ +connections: + bad: + provider: bitbucket + credential_env: SOME_TOKEN +""", + ) + + with pytest.raises(ProviderConfigError, match="unknown provider"): + load_registry(config_path=config, settings=mock_settings) + + +def test_rejects_explicit_connection_named_like_the_implicit_github_connection( + tmp_path, mock_settings, monkeypatch +): + """An explicit connection named 'github-default' would collide in the + adapter cache with the implicit zero-config GitHub connection, silently + mixing up which credential/host a resolution uses.""" + monkeypatch.setenv("SOME_TOKEN", "secret") + config = _write_config( + tmp_path, + """ +connections: + github-default: + provider: github + credential_env: SOME_TOKEN +""", + ) + + with pytest.raises(ProviderConfigError, match="reserved implicit connection name"): + load_registry(config_path=config, settings=mock_settings) + + +def test_rejects_repository_with_unknown_connection(tmp_path, mock_settings, monkeypatch): + monkeypatch.setenv("ACME_GITLAB_TOKEN", "secret") + config = _write_config( + tmp_path, + """ +connections: + acme-gitlab: + provider: gitlab + credential_env: ACME_GITLAB_TOKEN + +repositories: + payments-api: + provider: gitlab + connection: does-not-exist + namespace: platform/payments +""", + ) + + with pytest.raises(ProviderConfigError, match="unknown connection"): + load_registry(config_path=config, settings=mock_settings) + + +def test_rejects_connection_with_missing_credential_env_var(tmp_path, mock_settings, monkeypatch): + monkeypatch.delenv("UNSET_TOKEN_VAR", raising=False) + config = _write_config( + tmp_path, + """ +connections: + acme-gitlab: + provider: gitlab + credential_env: UNSET_TOKEN_VAR +""", + ) + + with pytest.raises(ProviderConfigError, match="UNSET_TOKEN_VAR"): + load_registry(config_path=config, settings=mock_settings) + + +def test_rejects_repository_namespace_excluded_by_allowed_namespaces( + tmp_path, mock_settings, monkeypatch +): + monkeypatch.setenv("ACME_GITLAB_TOKEN", "secret") + config = _write_config( + tmp_path, + """ +connections: + acme-gitlab: + provider: gitlab + credential_env: ACME_GITLAB_TOKEN + allowed_namespaces: ["platform/other"] + +repositories: + payments-api: + provider: gitlab + connection: acme-gitlab + namespace: platform/payments +""", + ) + + with pytest.raises(ProviderConfigError, match="allowed_namespaces"): + load_registry(config_path=config, settings=mock_settings) + + +def test_rejects_connection_with_string_allowed_namespaces(tmp_path, mock_settings, monkeypatch): + monkeypatch.setenv("ACME_GITLAB_TOKEN", "secret") + config = _write_config( + tmp_path, + """ +connections: + acme-gitlab: + provider: gitlab + credential_env: ACME_GITLAB_TOKEN + allowed_namespaces: "platform/payments" +""", + ) + + with pytest.raises(ProviderConfigError, match="allowed_namespaces.*list"): + load_registry(config_path=config, settings=mock_settings) + + +def test_rejects_connection_with_non_string_allowed_namespaces_elements( + tmp_path, mock_settings, monkeypatch +): + monkeypatch.setenv("ACME_GITLAB_TOKEN", "secret") + config = _write_config( + tmp_path, + """ +connections: + acme-gitlab: + provider: gitlab + credential_env: ACME_GITLAB_TOKEN + allowed_namespaces: [123, 456] +""", + ) + + with pytest.raises(ProviderConfigError, match="allowed_namespaces.*list of strings"): + load_registry(config_path=config, settings=mock_settings) + + +def test_rejects_repository_provider_mismatched_with_connection( + tmp_path, mock_settings, monkeypatch +): + monkeypatch.setenv("ACME_GITLAB_TOKEN", "secret") + config = _write_config( + tmp_path, + """ +connections: + acme-gitlab: + provider: gitlab + credential_env: ACME_GITLAB_TOKEN + +repositories: + payments-api: + provider: github + connection: acme-gitlab + namespace: platform/payments +""", + ) + + with pytest.raises(ProviderConfigError, match="provider"): + load_registry(config_path=config, settings=mock_settings) + + +def test_rejects_two_repositories_colliding_on_provider_and_namespace( + tmp_path, mock_settings, monkeypatch +): + monkeypatch.setenv("ACME_GITLAB_TOKEN", "secret") + config = _write_config( + tmp_path, + """ +connections: + acme-gitlab: + provider: gitlab + credential_env: ACME_GITLAB_TOKEN + +repositories: + payments-api: + provider: gitlab + connection: acme-gitlab + namespace: platform/payments + payments-api-dup: + provider: gitlab + connection: acme-gitlab + namespace: platform/payments +""", + ) + + with pytest.raises(ProviderConfigError, match="namespace"): + load_registry(config_path=config, settings=mock_settings) + + +def test_rejects_non_mapping_top_level_document(tmp_path, mock_settings): + config = _write_config(tmp_path, "- just\n- a\n- list\n") + + with pytest.raises(ProviderConfigError, match="mapping"): + load_registry(config_path=config, settings=mock_settings) + + +def test_rejects_non_mapping_connections_section(tmp_path, mock_settings): + config = _write_config(tmp_path, "connections: [not, a, mapping]\n") + + with pytest.raises(ProviderConfigError, match="'connections'"): + load_registry(config_path=config, settings=mock_settings) + + +def test_rejects_non_mapping_repositories_section(tmp_path, mock_settings): + config = _write_config(tmp_path, "repositories: [not, a, mapping]\n") + + with pytest.raises(ProviderConfigError, match="'repositories'"): + load_registry(config_path=config, settings=mock_settings) + + +def test_rejects_invalid_yaml_syntax(tmp_path, mock_settings): + config = _write_config(tmp_path, "connections: [unterminated\n") + + with pytest.raises(ProviderConfigError, match="invalid YAML"): + load_registry(config_path=config, settings=mock_settings) diff --git a/tests/unit/integrations/source_control/test_registry_resolve.py b/tests/unit/integrations/source_control/test_registry_resolve.py new file mode 100644 index 000000000..983b15954 --- /dev/null +++ b/tests/unit/integrations/source_control/test_registry_resolve.py @@ -0,0 +1,172 @@ +"""Tests for Registry.resolve() and adapter factory registration.""" + +from unittest.mock import AsyncMock + +import pytest +from pydantic import SecretStr + +from forge.integrations.source_control import registry as registry_module +from forge.integrations.source_control.contracts import Provider +from forge.integrations.source_control.errors import NotFoundError, ProviderConfigError +from forge.integrations.source_control.registry import ( + IMPLICIT_GITHUB_CONNECTION_NAME, + load_registry, + register_adapter_factory, +) + + +@pytest.fixture(autouse=True) +def _clean_adapter_factories(monkeypatch): + """Each test starts with no adapter factories registered.""" + monkeypatch.setattr(registry_module, "_ADAPTER_FACTORIES", {}) + + +def _write_config(tmp_path, content: str): + path = tmp_path / "repos.yaml" + path.write_text(content) + return path + + +def test_resolve_by_explicit_repository_id(tmp_path, mock_settings, monkeypatch): + monkeypatch.setenv("ACME_GITLAB_TOKEN", "secret") + config = _write_config( + tmp_path, + """ +connections: + acme-gitlab: + provider: gitlab + credential_env: ACME_GITLAB_TOKEN + +repositories: + payments-api: + provider: gitlab + connection: acme-gitlab + namespace: platform/payments + change_request_mode: direct +""", + ) + registry = load_registry(config_path=config, settings=mock_settings) + + resolved = registry.resolve("payments-api") + + assert resolved.repo_ref.id == "payments-api" + assert resolved.repo_ref.namespace == "platform/payments" + assert resolved.connection.name == "acme-gitlab" + + +def test_resolve_by_explicit_namespace_with_provider_hint(tmp_path, mock_settings, monkeypatch): + monkeypatch.setenv("ACME_GITLAB_TOKEN", "secret") + config = _write_config( + tmp_path, + """ +connections: + acme-gitlab: + provider: gitlab + credential_env: ACME_GITLAB_TOKEN + +repositories: + payments-api: + provider: gitlab + connection: acme-gitlab + namespace: platform/payments +""", + ) + registry = load_registry(config_path=config, settings=mock_settings) + + resolved = registry.resolve("platform/payments", provider_hint=Provider.GITLAB) + + assert resolved.repo_ref.id == "payments-api" + + +def test_resolve_falls_back_to_implicit_github_connection(tmp_path, mock_settings): + registry = load_registry(config_path=tmp_path / "missing.yaml", settings=mock_settings) + + resolved = registry.resolve("acme/payments") + + assert resolved.repo_ref.id == "acme/payments" + assert resolved.repo_ref.namespace == "acme/payments" + assert resolved.repo_ref.provider is Provider.GITHUB + assert resolved.repo_ref.change_request_mode == "fork" + assert resolved.connection.name == IMPLICIT_GITHUB_CONNECTION_NAME + assert resolved.connection.allowed_namespaces is None + + +def test_resolve_raises_not_found_for_provider_with_no_implicit_default(tmp_path, mock_settings): + registry = load_registry(config_path=tmp_path / "missing.yaml", settings=mock_settings) + + with pytest.raises(NotFoundError): + registry.resolve("platform/payments", provider_hint=Provider.GITLAB) + + +def test_resolve_raises_provider_config_error_when_implicit_credential_env_unset( + tmp_path, mock_settings +): + settings = mock_settings.model_copy(update={"github_token": SecretStr("")}) + registry = load_registry(config_path=tmp_path / "missing.yaml", settings=settings) + + with pytest.raises(ProviderConfigError, match="GITHUB_TOKEN"): + registry.resolve("acme/payments") + + +def test_resolve_uses_registered_adapter_factory(tmp_path, mock_settings): + registry = load_registry(config_path=tmp_path / "missing.yaml", settings=mock_settings) + sentinel_adapter = object() + register_adapter_factory(Provider.GITHUB, lambda _connection: sentinel_adapter) + + resolved = registry.resolve("acme/payments") + + assert resolved.adapter is sentinel_adapter + + +def test_resolve_adapter_is_none_when_no_factory_registered(tmp_path, mock_settings): + registry = load_registry(config_path=tmp_path / "missing.yaml", settings=mock_settings) + + resolved = registry.resolve("acme/payments") + + assert resolved.adapter is None + + +def test_resolve_caches_one_adapter_per_connection(tmp_path, mock_settings): + """Repeated resolve() calls against the same connection must reuse the + same adapter instance (and its underlying HTTP client), not construct a + fresh one -- constructing on every call leaks a connection pool per call.""" + registry = load_registry(config_path=tmp_path / "missing.yaml", settings=mock_settings) + factory_calls: list[object] = [] + register_adapter_factory( + Provider.GITHUB, lambda connection: factory_calls.append(connection) or object() + ) + + first = registry.resolve("acme/payments").adapter + second = registry.resolve("acme/payments").adapter + + assert first is second + assert len(factory_calls) == 1 + + +@pytest.mark.asyncio +async def test_aclose_closes_every_cached_adapter(tmp_path, mock_settings): + registry = load_registry(config_path=tmp_path / "missing.yaml", settings=mock_settings) + adapter = AsyncMock() + register_adapter_factory(Provider.GITHUB, lambda _connection: adapter) + registry.resolve("acme/payments") + + await registry.aclose() + + adapter.close.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_aclose_is_a_noop_when_nothing_was_ever_resolved(tmp_path, mock_settings): + registry = load_registry(config_path=tmp_path / "missing.yaml", settings=mock_settings) + + await registry.aclose() # must not raise + + +def test_get_registry_is_cached(): + registry_module.get_registry.cache_clear() + try: + first = registry_module.get_registry() + second = registry_module.get_registry() + assert first is second + finally: + registry_module.get_registry.cache_clear() diff --git a/tests/unit/models/test_events.py b/tests/unit/models/test_events.py index a0dd5eae8..bfdbae889 100644 --- a/tests/unit/models/test_events.py +++ b/tests/unit/models/test_events.py @@ -15,7 +15,7 @@ class TestEventSource: def test_event_sources_exist(self): """Verify event sources are defined.""" assert EventSource.JIRA.value == "jira" - assert EventSource.GITHUB.value == "github" + assert EventSource.SOURCE_CONTROL.value == "source_control" class TestEventStatus: @@ -52,12 +52,12 @@ def test_create_github_event(self): """Create a GitHub webhook event.""" event = WebhookEvent( event_id="evt-002", - source=EventSource.GITHUB, + source=EventSource.SOURCE_CONTROL, event_type="check_run", ticket_key="TEST-123", payload={"action": "completed"}, ) - assert event.source == EventSource.GITHUB + assert event.source == EventSource.SOURCE_CONTROL assert event.event_type == "check_run" def test_default_status_is_pending(self): @@ -130,7 +130,7 @@ def test_event_has_received_at(self): """Event has received_at timestamp.""" event = WebhookEvent( event_id="evt-008", - source=EventSource.GITHUB, + source=EventSource.SOURCE_CONTROL, event_type="pull_request", ticket_key="TEST-456", ) diff --git a/tests/unit/orchestrator/test_worker.py b/tests/unit/orchestrator/test_worker.py index c5e19caba..87772ad2e 100644 --- a/tests/unit/orchestrator/test_worker.py +++ b/tests/unit/orchestrator/test_worker.py @@ -1,18 +1,41 @@ """Unit tests for the orchestrator worker.""" +from datetime import UTC, datetime from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch import pytest +from forge.integrations.source_control.contracts import ( + Actor, + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + CheckStatus, + EventKind, + NormalizedEvent, + Provider, + RepositoryRef, + Review, + ReviewComment, + ReviewState, +) from forge.models.events import EventSource from forge.orchestrator.worker import ( OrchestratorWorker, - _cleanup_terminal_workspace, _has_new_reportable_error, _report_new_workflow_error, ) -from forge.queue.models import QueueMessage +from forge.queue.models import ( + QueueMessage, + normalized_event_to_dict, +) +from forge.workflow.utils.source_control import identity_for + + +def _patch_adapter(repo_ref: RepositoryRef, adapter): + """Patch worker.get_adapter to resolve to the given (repo_ref, adapter) pair.""" + return patch("forge.orchestrator.worker.get_adapter", return_value=(repo_ref, adapter)) @pytest.mark.parametrize( @@ -28,44 +51,6 @@ def test_has_new_reportable_error(result: dict, error_before_invoke: str | None, assert _has_new_reportable_error(result, error_before_invoke) is expected -@pytest.mark.asyncio -async def test_terminal_workflow_cleans_recreated_workspace(): - result = { - "ticket_key": "TEST-1", - "current_node": "complete", - "workspace_path": "/tmp/forge-TEST-1-repo", - "is_paused": False, - } - torn_down = { - **result, - "workspace_path": None, - "current_node": "workspace_complete", - } - - with patch( - "forge.orchestrator.worker.teardown_workspace", - AsyncMock(return_value=torn_down), - ) as teardown: - cleaned = await _cleanup_terminal_workspace(result) - - teardown.assert_awaited_once_with(result) - assert cleaned["workspace_path"] is None - assert cleaned["current_node"] == "complete" - - -@pytest.mark.asyncio -async def test_nonterminal_workflow_keeps_workspace(): - result = { - "current_node": "human_review_gate", - "workspace_path": "/tmp/forge-TEST-1-repo", - } - - with patch("forge.orchestrator.worker.teardown_workspace", AsyncMock()) as teardown: - assert await _cleanup_terminal_workspace(result) is result - - teardown.assert_not_awaited() - - @pytest.mark.asyncio async def test_report_new_workflow_error_posts_once(): result = { @@ -151,14 +136,14 @@ def _multi_repo_pr_state() -> dict: "current_pr_url": "https://github.com/acme/frontend/pull/20", "pr_merged": False, "pull_requests": { - "acme/backend": { + "acme/backend:10": { "repo": "acme/backend", "number": 10, "url": "https://github.com/acme/backend/pull/10", "merged": False, "ci_status": "pending", }, - "acme/frontend": { + "acme/frontend:20": { "repo": "acme/frontend", "number": 20, "url": "https://github.com/acme/frontend/pull/20", @@ -175,24 +160,30 @@ async def test_multi_repo_merge_waits_for_every_pr() -> None: state = _multi_repo_pr_state() def merge_message(repo: str, number: int) -> QueueMessage: + event = _make_normalized_event( + kind=EventKind.CR_MERGED, + repo_ref=_sc_repo_ref(repo), + change_request=_sc_change_request(repo, number, ChangeRequestState.MERGED), + ) return QueueMessage( message_id=f"msg-{number}", event_id=f"evt-{number}", - source=EventSource.GITHUB, - event_type="pull_request", + source=EventSource.SOURCE_CONTROL, + event_type="cr_merged", ticket_key="TEST-123", payload={ "action": "closed", "pull_request": {"merged": True, "number": number}, "repository": {"full_name": repo}, }, + normalized_event=normalized_event_to_dict(event), ) 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["pull_requests"]["acme/backend:10"]["merged"] is True + assert partial["pull_requests"]["acme/frontend:20"]["merged"] is False assert partial["pr_merged"] is False assert partial["is_paused"] is True @@ -205,16 +196,24 @@ def merge_message(repo: str, number: int) -> QueueMessage: @pytest.mark.asyncio async def test_multi_repo_ci_webhook_selects_earlier_pr_from_review_gate() -> None: worker = OrchestratorWorker(consumer_name="test-worker") + ci_payload = { + "check_suite": {"status": "completed", "pull_requests": [{"number": 10}]}, + "repository": {"full_name": "acme/backend"}, + } + event = _make_normalized_event( + kind=EventKind.CHECK_UPDATED, + repo_ref=_sc_repo_ref("acme/backend"), + change_request=_sc_change_request("acme/backend", 10), + raw=ci_payload, + ) message = QueueMessage( message_id="msg-ci", event_id="evt-ci", - source=EventSource.GITHUB, - event_type="check_suite", + source=EventSource.SOURCE_CONTROL, + event_type="check_updated", ticket_key="TEST-123", - payload={ - "check_suite": {"status": "completed", "pull_requests": [{"number": 10}]}, - "repository": {"full_name": "acme/backend"}, - }, + payload=ci_payload, + normalized_event=normalized_event_to_dict(event), ) result = await worker._handle_resume_event(message, _multi_repo_pr_state()) @@ -233,17 +232,26 @@ async def test_multi_repo_approval_uses_common_state_cleanup_path() -> None: state["last_error"] = "stale review failure" state["revision_requested"] = True state["feedback_comment"] = "old feedback" + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("acme/backend"), + change_request=_sc_change_request("acme/backend", 10), + # author="" reproduces the original payload's absent sender, so the + # self-comment guard is skipped without a network login lookup. + review=Review(id="", state=ReviewState.APPROVED, body="Looks good", author=""), + ) message = QueueMessage( message_id="msg-approved", event_id="evt-approved", - source=EventSource.GITHUB, - event_type="pull_request_review", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", payload={ "review": {"state": "approved", "body": "Looks good"}, "pull_request": {"number": 10}, "repository": {"full_name": "acme/backend"}, }, + normalized_event=normalized_event_to_dict(event), ) result = await worker._handle_resume_event(message, state) @@ -254,32 +262,38 @@ async def test_multi_repo_approval_uses_common_state_cleanup_path() -> 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" + assert result["pull_requests"]["acme/backend:10"]["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 +async def test_multi_repo_review_selects_earlier_pr() -> None: + mock_adapter = AsyncMock() + mock_adapter.get_review_comments_for_submission.return_value = [] worker = OrchestratorWorker(consumer_name="test-worker") state = _multi_repo_pr_state() state["current_node"] = "wait_for_ci_gate" + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("acme/backend"), + change_request=_sc_change_request("acme/backend", 10), + review=Review(id="5", state=ReviewState.CHANGES_REQUESTED, body="Fix backend", author=""), + ) message = QueueMessage( message_id="msg-review", event_id="evt-review", - source=EventSource.GITHUB, - event_type="pull_request_review", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", payload={ "review": {"id": 5, "state": "changes_requested", "body": "Fix backend"}, "pull_request": {"number": 10}, "repository": {"full_name": "acme/backend"}, }, + normalized_event=normalized_event_to_dict(event), ) - result = await worker._handle_resume_event(message, state) + with _patch_adapter(_sc_repo_ref("acme/backend"), mock_adapter): + result = await worker._handle_resume_event(message, state) assert result["current_repo"] == "acme/backend" assert result["current_pr_number"] == 10 @@ -1218,22 +1232,25 @@ def _ci_state(self, node: str) -> dict: } def _check_suite_message(self, conclusion: str = "failure") -> QueueMessage: + raw = { + "action": "completed", + "check_suite": { + "status": "completed", + "conclusion": conclusion, + "head_branch": "forge/aisos-701", + "pull_requests": [{"number": 52}], + }, + "repository": {"full_name": "forge-sdlc/forge"}, + } + event = _make_normalized_event(kind=EventKind.CHECK_UPDATED, raw=raw) return QueueMessage( message_id="1-0", event_id="test-ci-001", - source=EventSource.GITHUB, - event_type="check_suite", + source=EventSource.SOURCE_CONTROL, + event_type="check_updated", ticket_key="AISOS-701", - payload={ - "action": "completed", - "check_suite": { - "status": "completed", - "conclusion": conclusion, - "head_branch": "forge/aisos-701", - "pull_requests": [{"number": 52}], - }, - "repository": {"full_name": "forge-sdlc/forge"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) @pytest.mark.asyncio @@ -1259,16 +1276,22 @@ async def test_check_suite_recognized_at_ci_evaluator(self, worker): async def test_incomplete_check_suite_does_not_unpause_at_ci_evaluator(self, worker): """A check_suite with status=in_progress must not wake up the workflow.""" state = self._ci_state("ci_evaluator") + event = _make_normalized_event( + kind=EventKind.CHECK_UPDATED, + check_suite_status=CheckStatus.IN_PROGRESS, + raw={ + "check_suite": {"status": "in_progress", "conclusion": None}, + "repository": {"full_name": "forge-sdlc/forge"}, + }, + ) message = QueueMessage( message_id="1-0", event_id="test-ci-002", - source=EventSource.GITHUB, - event_type="check_suite", + source=EventSource.SOURCE_CONTROL, + event_type="check_updated", ticket_key="AISOS-701", - payload={ - "check_suite": {"status": "in_progress", "conclusion": None}, - "repository": {"full_name": "forge-sdlc/forge"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) result = await worker._handle_resume_event(message, state) @@ -1597,13 +1620,26 @@ async def test_ci_webhook_at_review_gate_sets_pending_ci_event(self, worker): "is_paused": True, "pending_ci_event": False, "context": {}, - "pull_requests": {"org/repo": {"number": 42}}, + "pull_requests": {"org/repo:42": {"number": 42}}, } + raw = { + "repository": {"full_name": "org/repo"}, + "check_suite": { + "status": "completed", + "pull_requests": [{"number": 42}], + }, + } + event = _make_normalized_event( + kind=EventKind.CHECK_UPDATED, + repo_ref=_sc_repo_ref("org/repo"), + change_request=_sc_change_request("org/repo", 42), + raw=raw, + ) message = QueueMessage( message_id="msg-1", event_id="evt-1", - source=EventSource.GITHUB, - event_type="check_suite.completed", + source=EventSource.SOURCE_CONTROL, + event_type="check_updated", ticket_key="TEST-1", payload={ "repository": {"full_name": "org/repo"}, @@ -1612,6 +1648,7 @@ async def test_ci_webhook_at_review_gate_sets_pending_ci_event(self, worker): "pull_requests": [{"number": 42}], }, }, + normalized_event=normalized_event_to_dict(event), ) result = await worker._handle_resume_event(message, current_state) @@ -1631,25 +1668,77 @@ async def test_ci_webhook_at_ci_evaluator_does_not_set_pending_ci_event(self, wo "pending_ci_event": False, "context": {}, } - message = QueueMessage( - message_id="msg-2", - event_id="evt-2", - source=EventSource.GITHUB, - event_type="check_suite.completed", - ticket_key="TEST-1", - payload={ + event = _make_normalized_event( + kind=EventKind.CHECK_UPDATED, + raw={ "check_suite": { "status": "completed", "pull_requests": [{"number": 42}], } }, ) + message = QueueMessage( + message_id="msg-2", + event_id="evt-2", + source=EventSource.SOURCE_CONTROL, + event_type="check_updated", + ticket_key="TEST-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) result = await worker._handle_resume_event(message, current_state) assert result.get("is_paused") is False assert result.get("pending_ci_event", False) is False # not set for ci_evaluator + @pytest.mark.asyncio + @patch("forge.orchestrator.worker.post_status_comment", new_callable=AsyncMock) + async def test_review_arriving_during_in_flight_ci_cycle_is_not_dropped( + self, _mock_post_comment + ): + """A PR review submitted while a CI webhook is still being evaluated at + human_review_gate must not be silently discarded — it should unpause and + record revision_requested/feedback_comment, and pending_ci_event must stay + set so the in-flight CI cycle still runs to completion.""" + mock_adapter = AsyncMock() + mock_adapter.get_review_thread_comments.return_value = [] + + worker = OrchestratorWorker(consumer_name="test-worker") + # State as left by the CI webhook that arrived first: unpaused, but still + # parked at human_review_gate with pending_ci_event set. + state = { + "ticket_key": "TEST-123", + "current_node": "human_review_gate", + "is_paused": False, + "pending_ci_event": True, + "context": {}, + } + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + review=Review( + id="", state=ReviewState.CHANGES_REQUESTED, body="Needs changes", author="" + ), + ) + message = QueueMessage( + message_id="msg-124", + event_id="evt-124", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", + ticket_key="TEST-123", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + + with _patch_adapter(_sc_repo_ref("owner/repo"), mock_adapter): + result = await worker._handle_resume_event(message, state) + + assert result["revision_requested"] is True + assert result["feedback_comment"] == "Needs changes" + assert result["pending_ci_event"] is True + class TestHandleResumeEventReviewGates: """Tests for resuming workflows from human_review_gate and review_response_gate.""" @@ -1657,16 +1746,17 @@ class TestHandleResumeEventReviewGates: @pytest.mark.asyncio async def test_forge_github_login_is_cached_per_worker(self): worker = OrchestratorWorker.__new__(OrchestratorWorker) - mock_github = AsyncMock() - mock_github.get_authenticated_user.return_value = {"login": "forge-bot"} + worker._forge_github_logins = {} + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) + repo_ref = _sc_repo_ref("owner/repo") - with patch("forge.orchestrator.worker.GitHubClient", return_value=mock_github): - first = await worker._get_forge_github_login() - second = await worker._get_forge_github_login() + with _patch_adapter(repo_ref, mock_adapter): + first = await worker._get_forge_github_login(repo_ref) + second = await worker._get_forge_github_login(repo_ref) assert first == second == "forge-bot" - mock_github.get_authenticated_user.assert_awaited_once() - mock_github.close.assert_awaited_once() + mock_adapter.get_authenticated_identity.assert_awaited_once_with(repo_ref) @pytest.mark.asyncio async def test_forge_authored_pr_review_does_not_resume_review_workflow(self): @@ -1680,23 +1770,21 @@ async def test_forge_authored_pr_review_does_not_resume_review_workflow(self): "is_paused": True, "context": {}, } + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + actor=Actor(login="forge-bot", is_bot=True), + review=Review(id="99", state=ReviewState.COMMENTED, body="", author="forge-bot"), + ) message = QueueMessage( message_id="msg-forge-review", event_id="evt-forge-review", - source=EventSource.GITHUB, - event_type="pull_request_review:submitted", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-236", - payload={ - "review": { - "id": 99, - "state": "commented", - "body": "", - "user": {"login": "forge-bot", "type": "Bot"}, - }, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - "sender": {"login": "forge-bot", "type": "Bot"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) with ( @@ -1705,13 +1793,13 @@ async def test_forge_authored_pr_review_does_not_resume_review_workflow(self): "_get_forge_github_login", new=AsyncMock(return_value="forge-bot"), ) as get_forge_login, - patch("forge.orchestrator.worker.GitHubClient") as github_client, + patch("forge.orchestrator.worker.get_adapter") as get_adapter_mock, ): result = await worker._handle_resume_event(message, state) assert result is state get_forge_login.assert_awaited_once() - github_client.assert_not_called() + get_adapter_mock.assert_not_called() @pytest.mark.asyncio async def test_inline_reply_resumes_only_its_contested_thread(self): @@ -1726,27 +1814,32 @@ async def test_inline_reply_resumes_only_its_contested_thread(self): ], "context": {}, } + event = _make_normalized_event( + kind=EventKind.COMMENT_CREATED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + actor=Actor(login="reviewer", is_bot=False), + comment=ReviewComment( + id="12", + body="Please make this change after all.", + author="reviewer", + path="src/file.py", + in_reply_to="11", + ), + ) message = QueueMessage( message_id="msg-thread-reply", event_id="evt-thread-reply", - source=EventSource.GITHUB, - event_type="pull_request_review_comment:created", + source=EventSource.SOURCE_CONTROL, + event_type="comment_created", ticket_key="TEST-233", - payload={ - "comment": { - "id": 12, - "in_reply_to_id": 11, - "body": "Please make this change after all.", - }, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - "sender": {"login": "reviewer"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) - mock_github = AsyncMock() - mock_github.get_authenticated_user.return_value = {"login": "forge-bot"} - with patch("forge.orchestrator.worker.GitHubClient", return_value=mock_github): + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) + with _patch_adapter(_sc_repo_ref("owner/repo"), mock_adapter): result = await worker._handle_resume_event(message, state) assert result["is_paused"] is False @@ -1756,6 +1849,12 @@ async def test_inline_reply_resumes_only_its_contested_thread(self): @pytest.mark.asyncio async def test_standalone_inline_comment_is_actionable_at_response_gate(self): + """A non-reply inline comment (no in_reply_to) at review_response_gate must + still be actionable — it does NOT silently fall through to an unchanged + state. This is the second (own-id) branch the typed cutover preserved: it + unpauses and requests revision using the comment's OWN id as the review + thread id, leaving contested_comments untouched (no reply target to clear). + """ worker = OrchestratorWorker(consumer_name="test-worker") state = { "ticket_key": "TEST-233", @@ -1764,41 +1863,58 @@ async def test_standalone_inline_comment_is_actionable_at_response_gate(self): "contested_comments": [{"thread_id": "thread-a", "comment_id": 10}], "context": {}, } + event = _make_normalized_event( + kind=EventKind.COMMENT_CREATED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + actor=Actor(login="reviewer", is_bot=False), + comment=ReviewComment( + id="30", + body="Please cover this edge case.", + author="reviewer", + path="src/file.py", + ), + ) message = QueueMessage( message_id="msg-new-thread", event_id="evt-new-thread", - source=EventSource.GITHUB, - event_type="pull_request_review_comment:created", + source=EventSource.SOURCE_CONTROL, + event_type="comment_created", ticket_key="TEST-233", - payload={ - "comment": {"id": 30, "body": "Please cover this edge case."}, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - "sender": {"login": "reviewer"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) - mock_github = AsyncMock() - mock_github.get_authenticated_user.return_value = {"login": "forge-bot"} + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) - with patch("forge.orchestrator.worker.GitHubClient", return_value=mock_github): + with _patch_adapter(_sc_repo_ref("owner/repo"), mock_adapter): result = await worker._handle_resume_event(message, state) + assert result["is_paused"] is False assert result["revision_requested"] is True assert result["feedback_comment"] == "Please cover this edge case." assert result["contested_comments"] == state["contested_comments"] + # The own-comment id (not a reply target) becomes the review thread id. + assert result["context"]["review_thread_comment_id"] == 30 @pytest.mark.asyncio @patch("forge.orchestrator.worker.post_status_comment", new_callable=AsyncMock) - @patch("forge.orchestrator.worker.GitHubClient") - async def test_pr_review_changes_requested_at_review_response_gate( - self, mock_github_client, _mock_post_comment - ): + async def test_pr_review_changes_requested_at_review_response_gate(self, _mock_post_comment): """changes_requested at review_response_gate unpauses and clears contested_comments.""" - mock_gh = AsyncMock() - mock_gh.get_pull_request_review_comments.return_value = [ - {"path": "src/file.py", "position": 10, "body": "Please fix this."} + mock_adapter = AsyncMock() + mock_adapter.get_review_thread_comments.return_value = [ + Review( + id="t1", + state=ReviewState.COMMENTED, + body="", + author="", + comments=[ + ReviewComment( + id="1", path="src/file.py", line=10, body="Please fix this.", author="" + ) + ], + ) ] - mock_github_client.return_value = mock_gh worker = OrchestratorWorker(consumer_name="test-worker") state = { @@ -1810,20 +1926,30 @@ async def test_pr_review_changes_requested_at_review_response_gate( ], "context": {}, } + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + review=Review( + id="", + state=ReviewState.CHANGES_REQUESTED, + body="PR needs some work", + author="", + ), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": {"state": "changes_requested", "body": "PR needs some work"}, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) - result = await worker._handle_resume_event(message, state) + repo_ref = _sc_repo_ref("owner/repo") + with _patch_adapter(repo_ref, mock_adapter): + result = await worker._handle_resume_event(message, state) assert result is not state assert result["is_paused"] is False @@ -1831,29 +1957,21 @@ async def test_pr_review_changes_requested_at_review_response_gate( assert result["contested_comments"] == [] assert "PR needs some work" in result["feedback_comment"] assert "src/file.py" in result["feedback_comment"] - mock_gh.get_pull_request_review_comments.assert_called_once_with("owner", "repo", 42) + mock_adapter.get_review_thread_comments.assert_called_once_with( + repo_ref, identity_for(repo_ref, 42) + ) @pytest.mark.asyncio @patch("forge.orchestrator.worker.post_status_comment", new_callable=AsyncMock) - @patch("forge.orchestrator.worker.GitHubClient") - async def test_pr_review_with_review_id_calls_get_review_comments( - self, mock_github_client, _mock_post_comment - ): - """When review payload contains a review ID, get_review_comments is called.""" - mock_gh = AsyncMock() - mock_gh.get_review_comments.return_value = [ - {"path": "src/file1.py", "position": 10, "body": "Fix position."}, - {"path": "src/file2.py", "line": 20, "body": "Fix line."}, - { - "path": "src/file2b.py", - "position": 4, - "line": 150, - "body": "Prefer the file line.", - }, - {"path": "src/file3.py", "original_line": 30, "body": "Fix original_line."}, - {"path": "src/file4.py", "body": "Fix none."}, + async def test_pr_review_with_review_id_calls_get_review_comments(self, _mock_post_comment): + """When review payload contains a review ID, get_review_comments_for_submission + is called (scoped to that review, not every unresolved thread).""" + mock_adapter = AsyncMock() + mock_adapter.get_review_comments_for_submission.return_value = [ + ReviewComment(id="1", path="src/file1.py", line=10, body="Fix line.", author=""), + ReviewComment(id="2", path="src/file2.py", line=20, body="Fix line.", author=""), + ReviewComment(id="3", path="src/file4.py", line=None, body="Fix none.", author=""), ] - mock_github_client.return_value = mock_gh worker = OrchestratorWorker(consumer_name="test-worker") state = { @@ -1865,24 +1983,30 @@ async def test_pr_review_with_review_id_calls_get_review_comments( ], "context": {}, } + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + review=Review( + id="999", + state=ReviewState.CHANGES_REQUESTED, + body="PR review body", + author="", + ), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": { - "id": 999, - "state": "changes_requested", - "body": "PR review body", - }, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) - result = await worker._handle_resume_event(message, state) + repo_ref = _sc_repo_ref("owner/repo") + with _patch_adapter(repo_ref, mock_adapter): + result = await worker._handle_resume_event(message, state) assert result is not state assert result["is_paused"] is False @@ -1892,28 +2016,49 @@ async def test_pr_review_with_review_id_calls_get_review_comments( assert "src/file1.py" in result["feedback_comment"] assert "(line 10)" in result["feedback_comment"] assert "(line 20)" in result["feedback_comment"] - assert "(line 150)" in result["feedback_comment"] - assert "(line 4)" not in result["feedback_comment"] - assert "(line 30)" in result["feedback_comment"] assert "(line ?)" in result["feedback_comment"] - mock_gh.get_review_comments.assert_called_once_with("owner", "repo", 42, 999) - mock_gh.get_pull_request_review_comments.assert_not_called() + mock_adapter.get_review_comments_for_submission.assert_called_once_with( + repo_ref, identity_for(repo_ref, 42), "999" + ) + mock_adapter.get_review_thread_comments.assert_not_called() @pytest.mark.asyncio @patch("forge.orchestrator.worker.post_status_comment", new_callable=AsyncMock) - @patch("forge.orchestrator.worker.GitHubClient") - async def test_pr_review_without_review_id_falls_back( - self, mock_github_client, _mock_post_comment - ): - """When review payload has NO review ID, get_pull_request_review_comments is called.""" - mock_gh = AsyncMock() - mock_gh.get_pull_request_review_comments.return_value = [ - {"path": "src/file1.py", "position": 10, "body": "Fix position."}, - {"path": "src/file2.py", "line": 20, "body": "Fix line."}, - {"path": "src/file3.py", "original_line": 30, "body": "Fix original_line."}, - {"path": "src/file4.py", "body": "Fix none."}, + async def test_pr_review_without_review_id_falls_back(self, _mock_post_comment): + """When review payload has NO review ID, get_review_thread_comments is + called (every unresolved thread) instead of the submission-scoped fetch.""" + mock_adapter = AsyncMock() + mock_adapter.get_review_thread_comments.return_value = [ + Review( + id="t1", + state=ReviewState.COMMENTED, + body="", + author="", + comments=[ + ReviewComment(id="1", path="src/file1.py", line=10, body="Fix line.", author="") + ], + ), + Review( + id="t2", + state=ReviewState.COMMENTED, + body="", + author="", + comments=[ + ReviewComment(id="2", path="src/file2.py", line=20, body="Fix line.", author="") + ], + ), + Review( + id="t3", + state=ReviewState.COMMENTED, + body="", + author="", + comments=[ + ReviewComment( + id="3", path="src/file4.py", line=None, body="Fix none.", author="" + ) + ], + ), ] - mock_github_client.return_value = mock_gh worker = OrchestratorWorker(consumer_name="test-worker") state = { @@ -1922,20 +2067,30 @@ async def test_pr_review_without_review_id_falls_back( "is_paused": True, "context": {}, } + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + review=Review( + id="", + state=ReviewState.CHANGES_REQUESTED, + body="PR review body", + author="", + ), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": {"state": "changes_requested", "body": "PR review body"}, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) - result = await worker._handle_resume_event(message, state) + repo_ref = _sc_repo_ref("owner/repo") + with _patch_adapter(repo_ref, mock_adapter): + result = await worker._handle_resume_event(message, state) assert result is not state assert result["is_paused"] is False @@ -1944,10 +2099,11 @@ async def test_pr_review_without_review_id_falls_back( assert "src/file1.py" in result["feedback_comment"] assert "(line 10)" in result["feedback_comment"] assert "(line 20)" in result["feedback_comment"] - assert "(line 30)" in result["feedback_comment"] assert "(line ?)" in result["feedback_comment"] - mock_gh.get_pull_request_review_comments.assert_called_once_with("owner", "repo", 42) - mock_gh.get_review_comments.assert_not_called() + mock_adapter.get_review_thread_comments.assert_called_once_with( + repo_ref, identity_for(repo_ref, 42) + ) + mock_adapter.get_review_comments_for_submission.assert_not_called() @pytest.mark.asyncio @patch("forge.orchestrator.worker.post_status_comment", new_callable=AsyncMock) @@ -1960,17 +2116,20 @@ async def test_pr_approve_at_review_response_gate(self, _mock_post_comment): "is_paused": True, "context": {}, } + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + review=Review(id="", state=ReviewState.APPROVED, body="Looks great!", author=""), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": {"state": "approved", "body": "Looks great!"}, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) result = await worker._handle_resume_event(message, state) @@ -1990,17 +2149,23 @@ async def test_pr_merge_at_review_response_gate(self, _mock_post_comment): "is_paused": True, "context": {}, } + event = _make_normalized_event( + kind=EventKind.CR_MERGED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42, ChangeRequestState.MERGED), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request", + source=EventSource.SOURCE_CONTROL, + event_type="cr_merged", ticket_key="TEST-123", payload={ "action": "closed", "pull_request": {"merged": True, "number": 42}, "repository": {"full_name": "owner/repo"}, }, + normalized_event=normalized_event_to_dict(event), ) result = await worker._handle_resume_event(message, state) @@ -2011,14 +2176,10 @@ async def test_pr_merge_at_review_response_gate(self, _mock_post_comment): @pytest.mark.asyncio @patch("forge.orchestrator.worker.post_status_comment", new_callable=AsyncMock) - @patch("forge.orchestrator.worker.GitHubClient") - async def test_pr_review_changes_requested_at_human_review_gate( - self, mock_github_client, _mock_post_comment - ): + async def test_pr_review_changes_requested_at_human_review_gate(self, _mock_post_comment): """changes_requested at human_review_gate unpauses and sets revision_requested.""" - mock_gh = AsyncMock() - mock_gh.get_pull_request_review_comments.return_value = [] - mock_github_client.return_value = mock_gh + mock_adapter = AsyncMock() + mock_adapter.get_review_thread_comments.return_value = [] worker = OrchestratorWorker(consumer_name="test-worker") state = { @@ -2027,20 +2188,26 @@ async def test_pr_review_changes_requested_at_human_review_gate( "is_paused": True, "context": {}, } + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + review=Review( + id="", state=ReviewState.CHANGES_REQUESTED, body="Needs changes", author="" + ), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": {"state": "changes_requested", "body": "Needs changes"}, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) - result = await worker._handle_resume_event(message, state) + with _patch_adapter(_sc_repo_ref("owner/repo"), mock_adapter): + result = await worker._handle_resume_event(message, state) assert result is not state assert result["is_paused"] is False @@ -2049,16 +2216,28 @@ async def test_pr_review_changes_requested_at_human_review_gate( @pytest.mark.asyncio @patch("forge.orchestrator.worker.post_status_comment", new_callable=AsyncMock) - @patch("forge.orchestrator.worker.GitHubClient") async def test_pr_commented_review_with_inline_at_review_response_gate( - self, mock_github_client, _mock_post_comment + self, _mock_post_comment ): """A 'commented' review with inline comments at review_response_gate is actionable.""" - mock_gh = AsyncMock() - mock_gh.get_pull_request_review_comments.return_value = [ - {"path": "src/app.py", "position": 5, "body": "Nit: rename this variable."} + mock_adapter = AsyncMock() + mock_adapter.get_review_thread_comments.return_value = [ + Review( + id="t1", + state=ReviewState.COMMENTED, + body="", + author="", + comments=[ + ReviewComment( + id="1", + path="src/app.py", + line=5, + body="Nit: rename this variable.", + author="", + ) + ], + ) ] - mock_github_client.return_value = mock_gh worker = OrchestratorWorker(consumer_name="test-worker") state = { @@ -2067,26 +2246,33 @@ async def test_pr_commented_review_with_inline_at_review_response_gate( "is_paused": True, "context": {}, } + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + review=Review(id="", state=ReviewState.COMMENTED, body="", author=""), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": {"state": "commented", "body": ""}, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) - result = await worker._handle_resume_event(message, state) + repo_ref = _sc_repo_ref("owner/repo") + with _patch_adapter(repo_ref, mock_adapter): + result = await worker._handle_resume_event(message, state) assert result is not state assert result["is_paused"] is False assert result["revision_requested"] is True assert "src/app.py" in result["feedback_comment"] - mock_gh.get_pull_request_review_comments.assert_called_once_with("owner", "repo", 42) + mock_adapter.get_review_thread_comments.assert_called_once_with( + repo_ref, identity_for(repo_ref, 42) + ) @pytest.mark.asyncio @patch("forge.orchestrator.worker.post_status_comment", new_callable=AsyncMock) @@ -2102,23 +2288,25 @@ async def test_pr_review_ignored_when_not_paused_at_review_response_gate( "is_paused": False, "context": {}, } + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + review=Review(id="", state=ReviewState.CHANGES_REQUESTED, body="Fix this", author=""), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": {"state": "changes_requested", "body": "Fix this"}, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) result = await worker._handle_resume_event(message, state) - assert result.get("revision_requested") is not True - assert result.get("feedback_comment") is None + assert result is state def test_review_response_gate_not_in_fresh_invoke_nodes(self): """review_response_gate must NOT use fresh-invoke — the gate re-pauses, @@ -2129,18 +2317,24 @@ def test_review_response_gate_not_in_fresh_invoke_nodes(self): @pytest.mark.asyncio @patch("forge.orchestrator.worker.post_status_comment", new_callable=AsyncMock) - @patch("forge.orchestrator.worker.GitHubClient") - async def test_review_response_gate_resume_routes_to_implement_review( - self, mock_github_client, _mock_post_comment - ): + async def test_review_response_gate_resume_routes_to_implement_review(self, _mock_post_comment): """After changes_requested at review_response_gate, state routes to implement_review.""" from forge.workflow.nodes.implement_review import route_review_response - mock_gh = AsyncMock() - mock_gh.get_pull_request_review_comments.return_value = [ - {"path": "src/main.py", "position": 7, "body": "Fix the typo here."} + mock_adapter = AsyncMock() + mock_adapter.get_review_thread_comments.return_value = [ + Review( + id="t1", + state=ReviewState.COMMENTED, + body="", + author="", + comments=[ + ReviewComment( + id="1", path="src/main.py", line=7, body="Fix the typo here.", author="" + ) + ], + ) ] - mock_github_client.return_value = mock_gh worker = OrchestratorWorker(consumer_name="test-worker") state = { @@ -2151,23 +2345,29 @@ async def test_review_response_gate_resume_routes_to_implement_review( "revision_requested": False, "context": {}, } + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + review=Review( + id="", + state=ReviewState.CHANGES_REQUESTED, + body="No, please apply the rename as requested", + author="", + ), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": { - "state": "changes_requested", - "body": "No, please apply the rename as requested", - }, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) - result = await worker._handle_resume_event(message, state) + with _patch_adapter(_sc_repo_ref("owner/repo"), mock_adapter): + result = await worker._handle_resume_event(message, state) assert route_review_response(result) == "implement_review" @@ -2188,36 +2388,37 @@ async def test_integration_bot_login_comment_without_prefix_processed_as_human_f "context": {}, } # Sender matches the bot login 'dev-user', but body does NOT contain the prefix signature + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + actor=Actor(login="dev-user", is_bot=False), + review=Review( + id="100", + state=ReviewState.CHANGES_REQUESTED, + body="!This is a human review comment without signature prefix.", + author="dev-user", + ), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review:submitted", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": { - "id": 100, - "state": "changes_requested", - "body": "!This is a human review comment without signature prefix.", - "user": {"login": "dev-user", "type": "User"}, - }, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - "sender": {"login": "dev-user", "type": "User"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) settings = MagicMock(forge_bot_comment_prefix="my-signature") + mock_adapter = AsyncMock() + mock_adapter.get_review_comments_for_submission.return_value = [] with ( patch.object(worker, "_get_forge_github_login", new=AsyncMock(return_value="dev-user")), patch("forge.orchestrator.worker.get_settings", return_value=settings), - patch("forge.orchestrator.worker.GitHubClient") as MockGH, + _patch_adapter(_sc_repo_ref("owner/repo"), mock_adapter), ): - mock_gh = AsyncMock() - mock_gh.get_review_comments.return_value = [] - MockGH.return_value = mock_gh - result = await worker._handle_resume_event(message, state) # It should be processed (not ignored), so state will have updated to resume (is_paused becomes False) @@ -2239,23 +2440,26 @@ async def test_integration_bot_login_comment_with_prefix_ignored_as_self_comment "context": {}, } # Sender matches the bot login 'dev-user', and body contains the prefix signature + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + actor=Actor(login="dev-user", is_bot=False), + review=Review( + id="100", + state=ReviewState.CHANGES_REQUESTED, + body="\n\nThis is an automated comment with signature.", + author="dev-user", + ), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review:submitted", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": { - "id": 100, - "state": "changes_requested", - "body": "\n\nThis is an automated comment with signature.", - "user": {"login": "dev-user", "type": "User"}, - }, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - "sender": {"login": "dev-user", "type": "User"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) settings = MagicMock(forge_bot_comment_prefix="my-signature") @@ -2263,11 +2467,7 @@ async def test_integration_bot_login_comment_with_prefix_ignored_as_self_comment with ( patch.object(worker, "_get_forge_github_login", new=AsyncMock(return_value="dev-user")), patch("forge.orchestrator.worker.get_settings", return_value=settings), - patch("forge.orchestrator.worker.GitHubClient") as MockGH, ): - mock_gh = AsyncMock() - MockGH.return_value = mock_gh - result = await worker._handle_resume_event(message, state) # It should be ignored (is_self_comment is True), so returns unchanged state @@ -2287,23 +2487,26 @@ async def test_integration_app_bot_comment_ending_in_bot_ignored_as_self_comment "context": {}, } # Sender is an App bot (ends with [bot]) and matches the bot login + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + actor=Actor(login="forge-bot[bot]", is_bot=True), + review=Review( + id="100", + state=ReviewState.CHANGES_REQUESTED, + body="Some comment body from app bot", + author="forge-bot[bot]", + ), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review:submitted", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": { - "id": 100, - "state": "changes_requested", - "body": "Some comment body from app bot", - "user": {"login": "forge-bot[bot]", "type": "Bot"}, - }, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - "sender": {"login": "forge-bot[bot]", "type": "Bot"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) settings = MagicMock(forge_bot_comment_prefix="my-signature") @@ -2313,11 +2516,7 @@ async def test_integration_app_bot_comment_ending_in_bot_ignored_as_self_comment worker, "_get_forge_github_login", new=AsyncMock(return_value="forge-bot") ), patch("forge.orchestrator.worker.get_settings", return_value=settings), - patch("forge.orchestrator.worker.GitHubClient") as MockGH, ): - mock_gh = AsyncMock() - MockGH.return_value = mock_gh - result = await worker._handle_resume_event(message, state) # It should be ignored because of the App bot suffix matching our bot login @@ -2337,38 +2536,39 @@ async def test_integration_other_app_bot_comment_ending_in_bot_is_not_ignored(se "context": {}, } # Sender is another App bot (ends with [bot], e.g., 'coderabbitai[bot]') + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + actor=Actor(login="coderabbitai[bot]", is_bot=True), + review=Review( + id="100", + state=ReviewState.CHANGES_REQUESTED, + body="!This is an external bot review comment.", + author="coderabbitai[bot]", + ), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review:submitted", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": { - "id": 100, - "state": "changes_requested", - "body": "!This is an external bot review comment.", - "user": {"login": "coderabbitai[bot]", "type": "Bot"}, - }, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - "sender": {"login": "coderabbitai[bot]", "type": "Bot"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) settings = MagicMock(forge_bot_comment_prefix="my-signature") + mock_adapter = AsyncMock() + mock_adapter.get_review_comments_for_submission.return_value = [] with ( patch.object( worker, "_get_forge_github_login", new=AsyncMock(return_value="forge-bot") ), patch("forge.orchestrator.worker.get_settings", return_value=settings), - patch("forge.orchestrator.worker.GitHubClient") as MockGH, + _patch_adapter(_sc_repo_ref("owner/repo"), mock_adapter), ): - mock_gh = AsyncMock() - mock_gh.get_review_comments.return_value = [] - MockGH.return_value = mock_gh - result = await worker._handle_resume_event(message, state) # It should be processed (not ignored) @@ -2390,23 +2590,26 @@ async def test_integration_legacy_fallback_no_prefix_ignored(self): "context": {}, } # Sender matches bot login, prefix is not configured + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + repo_ref=_sc_repo_ref("owner/repo"), + change_request=_sc_change_request("owner/repo", 42), + actor=Actor(login="dev-user", is_bot=False), + review=Review( + id="100", + state=ReviewState.CHANGES_REQUESTED, + body="Some body without signature", + author="dev-user", + ), + ) message = QueueMessage( message_id="msg-123", event_id="evt-123", - source=EventSource.GITHUB, - event_type="pull_request_review:submitted", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", ticket_key="TEST-123", - payload={ - "review": { - "id": 100, - "state": "changes_requested", - "body": "Some body without signature", - "user": {"login": "dev-user", "type": "User"}, - }, - "pull_request": {"number": 42}, - "repository": {"full_name": "owner/repo"}, - "sender": {"login": "dev-user", "type": "User"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) # Prefix is empty/None/disabled @@ -2415,13 +2618,633 @@ async def test_integration_legacy_fallback_no_prefix_ignored(self): with ( patch.object(worker, "_get_forge_github_login", new=AsyncMock(return_value="dev-user")), patch("forge.orchestrator.worker.get_settings", return_value=settings), - patch("forge.orchestrator.worker.GitHubClient") as MockGH, ): - mock_gh = AsyncMock() - MockGH.return_value = mock_gh - result = await worker._handle_resume_event(message, state) # It should be ignored under the legacy fallback because prefix is empty assert result is state assert result.get("is_paused") is True + + +def _make_normalized_event(**overrides) -> NormalizedEvent: + repo_ref = RepositoryRef( + id="acme/payments", + provider=Provider.GITHUB, + connection="default-github", + namespace="acme/payments", + default_branch="main", + change_request_mode="fork", + ) + change_request = ChangeRequest( + identity=ChangeRequestIdentity( + connection="default-github", repository_id="acme/payments", native_id=42 + ), + url="https://github.com/acme/payments/pull/42", + title="t", + body="", + state=ChangeRequestState.OPEN, + source_branch="feature", + target_branch="main", + draft=False, + ) + defaults = { + "id": "delivery-1", + "kind": EventKind.CR_OPENED, + "repo_ref": repo_ref, + "actor": Actor(login="octocat", is_bot=False), + "received_at": datetime(2026, 1, 1, tzinfo=UTC), + "change_request": change_request, + "raw": {}, + } + defaults.update(overrides) + return NormalizedEvent(**defaults) + + +def _sc_repo_ref(namespace: str = "owner/repo") -> RepositoryRef: + """Build a RepositoryRef for a given owner/repo namespace.""" + return RepositoryRef( + id=namespace, + provider=Provider.GITHUB, + connection="default-github", + namespace=namespace, + default_branch="main", + change_request_mode="fork", + ) + + +def _sc_change_request( + namespace: str = "owner/repo", + number: int = 42, + state: ChangeRequestState = ChangeRequestState.OPEN, +) -> ChangeRequest: + """Build a ChangeRequest carrying the PR number in native_id.""" + return ChangeRequest( + identity=ChangeRequestIdentity( + connection="default-github", repository_id=namespace, native_id=number + ), + url=f"https://github.com/{namespace}/pull/{number}", + title="t", + body="", + state=state, + source_branch="feature", + target_branch="main", + draft=False, + ) + + +class TestDeserializeEvent: + """Tests for NormalizedEvent reconstruction from a queue message.""" + + @pytest.fixture + def worker(self) -> OrchestratorWorker: + """Create a worker instance for testing.""" + return OrchestratorWorker(consumer_name="test-worker") + + def test_returns_none_for_jira_message(self, worker): + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.JIRA, + event_type="issue_updated", + ticket_key="PROJ-1", + payload={}, + ) + assert worker._deserialize_event(message) is None + + def test_deserializes_source_control_message(self, worker): + event = _make_normalized_event() + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="cr_opened", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + restored = worker._deserialize_event(message) + assert restored is not None + assert restored.kind == EventKind.CR_OPENED + assert restored.repo_ref.namespace == "acme/payments" + + +class TestIsPrdSpecPrEvent: + """Tests for PRD/spec proposals-PR detection off typed event fields.""" + + @pytest.fixture + def worker(self) -> OrchestratorWorker: + """Create a worker instance for testing.""" + return OrchestratorWorker(consumer_name="test-worker") + + def test_is_prd_pr_event_matches_by_repo_and_number(self, worker): + event = _make_normalized_event() + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="cr_updated", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = {"prd_pr_number": 42, "prd_pr_repo": "acme/payments"} + + assert worker._is_prd_pr_event(message, current_state) is True + + def test_is_prd_pr_event_false_when_number_differs(self, worker): + event = _make_normalized_event() + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="cr_updated", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = {"prd_pr_number": 99, "prd_pr_repo": "acme/payments"} + + assert worker._is_prd_pr_event(message, current_state) is False + + def test_is_prd_pr_event_false_for_jira_source(self, worker): + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.JIRA, + event_type="issue_updated", + ticket_key="PROJ-1", + payload={}, + ) + current_state = {"prd_pr_number": 42, "prd_pr_repo": "acme/payments"} + + assert worker._is_prd_pr_event(message, current_state) is False + + +class TestCiWebhookDetectionTypedFields: + """CI-webhook detection reads typed NormalizedEvent fields (Task 14).""" + + @pytest.fixture + def worker(self) -> OrchestratorWorker: + return OrchestratorWorker(consumer_name="test-worker") + + @pytest.mark.asyncio + async def test_check_run_completed_wakes_ci_evaluator(self, worker): + event = _make_normalized_event(kind=EventKind.CHECK_UPDATED) + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="check_updated", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "ci_evaluator", + "is_paused": True, + "pull_requests": {"acme/payments:42": {"number": 42, "repo": "acme/payments"}}, + "current_repo": "acme/payments", + "current_pr_number": 42, + } + + updated = await worker._handle_resume_event(message, current_state) + + assert updated["is_paused"] is False + + @pytest.mark.asyncio + async def test_incomplete_check_suite_does_not_wake_ci_evaluator(self, worker): + """A CHECK_UPDATED event whose suite is still in_progress must not unpause. + + The suite-completion nuance from the original truth table is preserved by + reading the normalized check_suite_status field (which survives the queue hop). + """ + event = _make_normalized_event( + kind=EventKind.CHECK_UPDATED, + check_suite_status=CheckStatus.IN_PROGRESS, + raw={"check_suite": {"status": "in_progress"}}, + ) + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="check_updated", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = {"current_node": "ci_evaluator", "is_paused": True, "context": {}} + + updated = await worker._handle_resume_event(message, current_state) + + assert updated is current_state + + @pytest.mark.asyncio + async def test_synchronize_push_event_wakes_ci_evaluator(self, worker): + """Preserve the original 'extra CI branch': a non-check, non-comment, + non-review, non-merged event with targets_implementation_pr=True still + wakes the workflow at ci_evaluator (branch (b) of the original truth table). + """ + event = _make_normalized_event(kind=EventKind.CR_UPDATED) # synchronize + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="cr_updated", + ticket_key="PROJ-1", + payload={ + "repository": {"full_name": "acme/payments"}, + "pull_request": {"number": 42}, + }, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "ci_evaluator", + "is_paused": True, + "context": {}, + "pull_requests": {"acme/payments:42": {"number": 42, "repo": "acme/payments"}}, + } + + updated = await worker._handle_resume_event(message, current_state) + + assert updated["is_paused"] is False + + @pytest.mark.asyncio + async def test_merged_pr_event_does_not_wake_ci_evaluator(self, worker): + """A merged change request is excluded from the 'extra CI branch' — it + must not set is_ci_webhook (matching the original merged-PR exclusion). + + This isolates the CI-branch exclusion: the event does NOT target an + implementation PR, so the (separately tested) PR-merge-at-review-gate + block does not fire and the paused ci_evaluator stays paused because no + CI signal was recognised. + """ + event = _make_normalized_event(kind=EventKind.CR_MERGED) + event.change_request.state = ChangeRequestState.MERGED + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="cr_merged", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "ci_evaluator", + "is_paused": True, + "context": {}, + } + + updated = await worker._handle_resume_event(message, current_state) + + # No CI signal recognised — the paused gate is not woken. + assert updated["is_paused"] is True + + @pytest.mark.asyncio + async def test_non_command_comment_does_not_set_ci_webhook(self, worker): + """A plain (non-command) COMMENT_CREATED at ci_evaluator must NOT fire the + CI-webhook branch — comment-kind events are excluded from branch (b). + """ + event = _make_normalized_event(kind=EventKind.COMMENT_CREATED) + event.comment = ReviewComment(id="1", body="just a regular comment", author="octocat") + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="comment_created", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = {"current_node": "ci_evaluator", "is_paused": True, "context": {}} + + updated = await worker._handle_resume_event(message, current_state) + + # CI-webhook branch did not fire — the paused gate stays paused. + assert updated["is_paused"] is True + + @pytest.mark.asyncio + async def test_review_submitted_does_not_set_ci_webhook(self, worker): + """A REVIEW_SUBMITTED event at ci_evaluator must NOT fire the CI-webhook + branch — review-kind events are excluded from branch (b). + """ + event = _make_normalized_event(kind=EventKind.REVIEW_SUBMITTED) + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = {"current_node": "ci_evaluator", "is_paused": True, "context": {}} + + updated = await worker._handle_resume_event(message, current_state) + + assert updated["is_paused"] is True + + +class TestSkipGateCommandTypedFields: + """skip-gate/unskip-gate/rebase detection reads typed fields (Task 14).""" + + @pytest.fixture + def worker(self) -> OrchestratorWorker: + return OrchestratorWorker(consumer_name="test-worker") + + @pytest.mark.asyncio + async def test_skip_gate_command_adds_check_name(self, worker): + event = _make_normalized_event(kind=EventKind.COMMENT_CREATED) + event.comment = ReviewComment(id="1", body="/forge skip-gate flaky-test", author="octocat") + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="comment_created", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = {"current_node": "ci_evaluator", "is_paused": True} + + with patch.object(worker, "_post_skip_gate_feedback", AsyncMock()): + updated = await worker._handle_resume_event(message, current_state) + + assert "flaky-test" in updated["ci_skipped_checks"] + assert updated["current_node"] == "ci_evaluator" + + @pytest.mark.asyncio + async def test_skip_gate_passes_typed_pr_and_sender_to_feedback(self, worker): + """pr_number, owner/repo and sender come from typed fields, not the payload.""" + event = _make_normalized_event(kind=EventKind.COMMENT_CREATED) + event.comment = ReviewComment(id="1", body="/forge skip-gate flaky-test", author="octocat") + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="comment_created", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = {"current_node": "ci_evaluator", "is_paused": True} + feedback = AsyncMock() + + with patch.object(worker, "_post_skip_gate_feedback", feedback): + await worker._handle_resume_event(message, current_state) + + feedback.assert_called_once() + kwargs = feedback.call_args.kwargs + assert kwargs["repo_ref"].namespace == "acme/payments" + assert kwargs["pr_number"] == 42 + assert kwargs["sender"] == "octocat" + + @pytest.mark.asyncio + async def test_rebase_command_routes_to_rebase_pr(self, worker): + """/forge rebase reads typed fields and routes to rebase_pr.""" + event = _make_normalized_event(kind=EventKind.COMMENT_CREATED) + event.comment = ReviewComment(id="1", body="/forge rebase", author="octocat") + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="comment_created", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "human_review_gate", + "is_paused": True, + "current_pr_number": 42, + } + feedback = AsyncMock() + + with patch.object(worker, "_post_rebase_feedback", feedback): + updated = await worker._handle_resume_event(message, current_state) + + assert updated["current_node"] == "rebase_pr" + assert updated["is_paused"] is False + assert updated["rebase_return_node"] == "human_review_gate" + feedback.assert_called_once() + kwargs = feedback.call_args.kwargs + assert kwargs["repo_ref"].namespace == "acme/payments" + assert kwargs["pr_number"] == 42 + assert kwargs["sender"] == "octocat" + + +class TestInlineReviewReplyTypedFields: + """Inline review-reply detection at review_response_gate reads typed fields.""" + + @pytest.fixture + def worker(self) -> OrchestratorWorker: + return OrchestratorWorker(consumer_name="test-worker") + + @pytest.mark.asyncio + async def test_inline_reply_clears_matching_contested_comment(self, worker): + # A pull_request_review_comment reply maps to COMMENT_CREATED with a path + # set and in_reply_to carrying the parent comment id (a str). The parent + # id is coerced to int to match the int comment ids persisted in state. + event = _make_normalized_event( + kind=EventKind.COMMENT_CREATED, + comment=ReviewComment( + id="2", body="fixed", author="octocat", path="src/x.py", in_reply_to="1" + ), + ) + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="comment_created", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "review_response_gate", + "is_paused": True, + "contested_comments": [{"comment_id": 1}], + } + + with patch.object(worker, "_get_forge_github_login", AsyncMock(return_value="forge-bot")): + updated = await worker._handle_resume_event(message, current_state) + + assert updated["revision_requested"] is True + assert updated["contested_comments"] == [] + assert updated["context"]["review_thread_comment_id"] == 1 + + @pytest.mark.asyncio + async def test_non_reply_inline_comment_is_still_actionable(self, worker): + """The two-branch question resolved: a non-reply inline comment (no + in_reply_to) at review_response_gate does NOT fall through to an unchanged + state (which would be a silent regression). It is handled by the preserved + second branch — unpause + revision using the comment's own id — leaving + contested threads untouched. Behavior is therefore NOT equivalent to a + fall-through, so both branches are kept. + """ + event = _make_normalized_event( + kind=EventKind.COMMENT_CREATED, + comment=ReviewComment( + id="30", body="Please cover this case.", author="octocat", path="src/x.py" + ), + ) + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="comment_created", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "review_response_gate", + "is_paused": True, + "contested_comments": [{"comment_id": 1}], + } + + with patch.object(worker, "_get_forge_github_login", AsyncMock(return_value="forge-bot")): + updated = await worker._handle_resume_event(message, current_state) + + assert updated is not current_state + assert updated["is_paused"] is False + assert updated["revision_requested"] is True + assert updated["feedback_comment"] == "Please cover this case." + # Non-reply: contested threads are preserved and the own id is recorded. + assert updated["contested_comments"] == [{"comment_id": 1}] + assert updated["context"]["review_thread_comment_id"] == 30 + + @pytest.mark.asyncio + async def test_top_level_issue_comment_does_not_match_this_block(self, worker): + """An issue comment (COMMENT_CREATED with no path) at review_response_gate + must NOT be treated as an inline review reply — matching the original + 'pull_request_review_comment' event-type restriction. With no other + review_response_gate handler, a paused gate stays paused. + """ + event = _make_normalized_event( + kind=EventKind.COMMENT_CREATED, + comment=ReviewComment(id="7", body="just a comment", author="octocat"), + ) + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="comment_created", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "review_response_gate", + "is_paused": True, + "context": {}, + } + + updated = await worker._handle_resume_event(message, current_state) + + assert updated is current_state + + +class TestHumanReviewGateTypedFields: + """Human-review-gate PR-review + PR-merge detection reads typed fields.""" + + @pytest.fixture + def worker(self) -> OrchestratorWorker: + return OrchestratorWorker(consumer_name="test-worker") + + @pytest.mark.asyncio + async def test_review_approved_sets_implementation_pr_approved(self, worker): + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + review=Review(id="1", state=ReviewState.APPROVED, body="", author="reviewer1"), + ) + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", + ticket_key="PROJ-1", + payload={ + "repository": {"full_name": "acme/payments"}, + "pull_request": {"number": 42}, + }, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "human_review_gate", + "is_paused": True, + "pull_requests": {"acme/payments:42": {"number": 42, "repo": "acme/payments"}}, + "current_repo": "acme/payments", + "current_pr_number": 42, + } + + with patch.object(worker, "_get_forge_github_login", AsyncMock(return_value="forge-bot")): + updated = await worker._handle_resume_event(message, current_state) + + assert updated["human_review_status"] == "approved" + + @pytest.mark.asyncio + async def test_pr_merged_at_review_gate_sets_pr_merged(self, worker): + event = _make_normalized_event(kind=EventKind.CR_MERGED) + event.change_request.state = ChangeRequestState.MERGED + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="cr_merged", + ticket_key="PROJ-1", + payload={ + "repository": {"full_name": "acme/payments"}, + "pull_request": {"merged": True, "number": 42}, + }, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "human_review_gate", + "is_paused": True, + "pull_requests": {"acme/payments:42": {"number": 42, "repo": "acme/payments"}}, + "current_repo": "acme/payments", + "current_pr_number": 42, + } + + updated = await worker._handle_resume_event(message, current_state) + + assert updated.get("pr_merged") is True + + @pytest.mark.asyncio + async def test_dismissed_review_does_not_trigger_revision(self, worker): + """A dismissed review (an admin unblocking a stale review) must not be + mistaken for an active COMMENTED review requesting changes -- it maps + to its own ReviewState.DISMISSED, which matches neither the APPROVED + nor the (CHANGES_REQUESTED, COMMENTED) branches, so state is left + unchanged, same as the original raw-string-based behavior.""" + event = _make_normalized_event( + kind=EventKind.REVIEW_SUBMITTED, + review=Review(id="1", state=ReviewState.DISMISSED, body="", author="reviewer1"), + ) + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", + ticket_key="PROJ-1", + payload={ + "repository": {"full_name": "acme/payments"}, + "pull_request": {"number": 42}, + }, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "human_review_gate", + "is_paused": True, + "pull_requests": {"acme/payments:42": {"number": 42, "repo": "acme/payments"}}, + "current_repo": "acme/payments", + "current_pr_number": 42, + } + + with patch.object(worker, "_get_forge_github_login", AsyncMock(return_value="forge-bot")): + updated = await worker._handle_resume_event(message, current_state) + + assert "human_review_status" not in updated + assert updated.get("revision_requested") is not True + assert updated.get("is_paused") is True diff --git a/tests/unit/orchestrator/test_worker_prd_pr.py b/tests/unit/orchestrator/test_worker_prd_pr.py index 9f523198f..d86c18993 100644 --- a/tests/unit/orchestrator/test_worker_prd_pr.py +++ b/tests/unit/orchestrator/test_worker_prd_pr.py @@ -1,24 +1,212 @@ """Tests for PRD PR event handling in the worker.""" +from datetime import UTC, datetime from unittest.mock import AsyncMock, MagicMock, patch import pytest +from forge.integrations.source_control.contracts import ( + Actor, + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + EventKind, + NormalizedEvent, + Provider, + RepositoryRef, + Review, + ReviewComment, + ReviewState, +) from forge.models.events import EventSource from forge.orchestrator.worker import OrchestratorWorker -from forge.queue.models import QueueMessage +from forge.queue.models import QueueMessage, normalized_event_to_dict from forge.workflow.utils.automated_review_triage import AutomatedReviewDecision +from forge.workflow.utils.source_control import identity_for + + +def _repo_ref_for(namespace: str) -> RepositoryRef: + return RepositoryRef( + id=namespace, + provider=Provider.GITHUB, + connection="default-github", + namespace=namespace, + default_branch="main", + change_request_mode="fork", + ) + + +def _patch_adapter(repo_ref: RepositoryRef, adapter): + """Patch worker.get_adapter to resolve to the given (repo_ref, adapter) pair.""" + return patch("forge.orchestrator.worker.get_adapter", return_value=(repo_ref, adapter)) + + +_REVIEW_STATES = { + "approved": ReviewState.APPROVED, + "changes_requested": ReviewState.CHANGES_REQUESTED, + "commented": ReviewState.COMMENTED, + "pending": ReviewState.PENDING, +} + + +def _normalized_from_payload(event_type: str, payload: dict) -> NormalizedEvent: + """Build the NormalizedEvent a GitHub webhook payload would produce. + + Mirrors GitHubAdapter.parse_webhook for the event types exercised here so the + typed detection in _handle_resume_event runs against realistic data while the + raw payload is still carried for the triage blocks that read it. + """ + base_type = event_type.split(":", 1)[0] + repo_full = payload.get("repository", {}).get("full_name", "") + sender = payload.get("sender", {}) + sender_login = sender.get("login", "") + actor = Actor( + login=sender_login, + is_bot=sender.get("type") == "Bot" or "[bot]" in sender_login, + ) + repo_ref = RepositoryRef( + id=repo_full, + provider=Provider.GITHUB, + connection="default-github", + namespace=repo_full, + default_branch="main", + change_request_mode="fork", + ) + + change_request = None + pr = payload.get("pull_request") + issue = payload.get("issue") + if pr is not None: + if pr.get("merged", False): + state = ChangeRequestState.MERGED + elif pr.get("state") == "closed": + state = ChangeRequestState.CLOSED + else: + state = ChangeRequestState.OPEN + change_request = ChangeRequest( + identity=ChangeRequestIdentity( + connection="default-github", + repository_id=repo_full, + native_id=pr.get("number"), + ), + url=pr.get("html_url", ""), + title=pr.get("title", ""), + body=pr.get("body", "") or "", + state=state, + source_branch="", + target_branch="", + draft=False, + ) + elif issue is not None: + change_request = ChangeRequest( + identity=ChangeRequestIdentity( + connection="default-github", + repository_id=repo_full, + native_id=issue.get("number"), + ), + url=issue.get("html_url", ""), + title=issue.get("title", ""), + body=issue.get("body", "") or "", + state=ChangeRequestState.OPEN, + source_branch="", + target_branch="", + draft=False, + ) + + comment = None + review = None + if base_type in ("issue_comment", "pull_request_review_comment"): + kind = EventKind.COMMENT_CREATED + raw_comment = payload.get("comment", {}) + in_reply = raw_comment.get("in_reply_to_id") + comment = ReviewComment( + id=str(raw_comment.get("id", "")), + body=raw_comment.get("body", "") or "", + author=(raw_comment.get("user") or {}).get("login", ""), + path=raw_comment.get("path"), + line=raw_comment.get("line"), + in_reply_to=str(in_reply) if in_reply is not None else None, + ) + elif base_type == "pull_request_review": + kind = EventKind.REVIEW_SUBMITTED + raw_review = payload.get("review", {}) + review = Review( + id=str(raw_review.get("id", "")), + state=_REVIEW_STATES.get( + (raw_review.get("state") or "").lower(), ReviewState.COMMENTED + ), + body=raw_review.get("body", "") or "", + author=(raw_review.get("user") or {}).get("login", ""), + comments=[], + ) + elif base_type == "pull_request": + if change_request is not None and change_request.state == ChangeRequestState.MERGED: + kind = EventKind.CR_MERGED + elif change_request is not None and change_request.state == ChangeRequestState.CLOSED: + kind = EventKind.CR_CLOSED + else: + kind = EventKind.CR_UPDATED + else: + kind = EventKind.CR_UPDATED + + return NormalizedEvent( + id="evt-1", + kind=kind, + repo_ref=repo_ref, + actor=actor, + received_at=datetime(2026, 1, 1, tzinfo=UTC), + change_request=change_request, + comment=comment, + review=review, + raw=payload, + ) def _make_message(event_type: str, payload: dict, ticket_key: str = "TEST-123") -> QueueMessage: return QueueMessage( message_id="msg-1", event_id="evt-1", - source=EventSource.GITHUB, + source=EventSource.SOURCE_CONTROL, event_type=event_type, ticket_key=ticket_key, payload=payload, + normalized_event=normalized_event_to_dict(_normalized_from_payload(event_type, payload)), + ) + + +def _make_normalized_event(**overrides) -> NormalizedEvent: + """A canned NormalizedEvent (repo acme/payments, PR 42) for typed-field tests.""" + repo_ref = RepositoryRef( + id="acme/payments", + provider=Provider.GITHUB, + connection="default-github", + namespace="acme/payments", + default_branch="main", + change_request_mode="fork", + ) + change_request = ChangeRequest( + identity=ChangeRequestIdentity( + connection="default-github", repository_id="acme/payments", native_id=42 + ), + url="https://github.com/acme/payments/pull/42", + title="t", + body="", + state=ChangeRequestState.OPEN, + source_branch="feature", + target_branch="main", + draft=False, ) + defaults = { + "id": "delivery-1", + "kind": EventKind.CR_OPENED, + "repo_ref": repo_ref, + "actor": Actor(login="octocat", is_bot=False), + "received_at": datetime(2026, 1, 1, tzinfo=UTC), + "change_request": change_request, + "raw": {}, + } + defaults.update(overrides) + return NormalizedEvent(**defaults) def _prd_gate_state(**overrides) -> dict: @@ -48,6 +236,7 @@ def worker(): w = OrchestratorWorker.__new__(OrchestratorWorker) w._post_terminal_error_comment = AsyncMock() w._post_resume_ack_comment = AsyncMock() + w._forge_github_logins = {} return w @@ -182,18 +371,19 @@ async def test_changes_requested_sets_feedback(self, worker): ) state = _prd_gate_state() - with patch("forge.orchestrator.worker.GitHubClient") as MockGH: - mock_gh = MagicMock() - mock_gh.get_pull_request_review_threads = AsyncMock(return_value=[]) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh + repo_ref = _repo_ref_for("org/proposals") + mock_adapter = AsyncMock() + mock_adapter.get_review_thread_comments.return_value = [] + with _patch_adapter(repo_ref, mock_adapter): result = await worker._handle_resume_event(msg, state) assert result["is_paused"] is False assert result["revision_requested"] is True assert "more detail" in result["feedback_comment"] - mock_gh.get_pull_request_review_threads.assert_called_once_with("org", "proposals", 7) + mock_adapter.get_review_thread_comments.assert_called_once_with( + repo_ref, identity_for(repo_ref, 7) + ) @pytest.mark.asyncio async def test_approved_review_is_ignored(self, worker): @@ -224,18 +414,28 @@ async def test_mixed_threads_revise_accepts_and_reply_to_contested(self, worker) }, ) threads = [ - { - "thread_id": "accept-thread", - "path": "prd.md", - "line": 10, - "comments": [{"comment_id": 10, "body": "Clarify authorization."}], - }, - { - "thread_id": "reply-thread", - "path": "prd.md", - "line": 20, - "comments": [{"comment_id": 20, "body": "Rename the product."}], - }, + Review( + id="accept-thread", + state=ReviewState.COMMENTED, + body="", + author="", + comments=[ + ReviewComment( + id="10", path="prd.md", line=10, body="Clarify authorization.", author="" + ) + ], + ), + Review( + id="reply-thread", + state=ReviewState.COMMENTED, + body="", + author="", + comments=[ + ReviewComment( + id="20", path="prd.md", line=20, body="Rename the product.", author="" + ) + ], + ), ] decisions = [ { @@ -257,8 +457,11 @@ async def test_mixed_threads_revise_accepts_and_reply_to_contested(self, worker) ] state = _prd_gate_state(prd_content="# Current PRD") + mock_adapter = AsyncMock() + mock_adapter.get_review_thread_comments.return_value = threads + with ( - patch("forge.orchestrator.worker.GitHubClient") as MockGH, + _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter), patch( "forge.orchestrator.worker.triage_proposal_review_threads", new=AsyncMock(return_value=decisions), @@ -268,10 +471,6 @@ async def test_mixed_threads_revise_accepts_and_reply_to_contested(self, worker) new=AsyncMock(), ) as reply_decisions, ): - mock_gh = MagicMock() - mock_gh.get_pull_request_review_threads = AsyncMock(return_value=threads) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh result = await worker._handle_resume_event(msg, state) assert result["revision_requested"] is True @@ -300,12 +499,9 @@ async def test_comment_sets_feedback(self, worker): automated_review_revision_pending=True, ) - with patch("forge.orchestrator.worker.GitHubClient") as MockGH: - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh - + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) + with _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter): result = await worker._handle_resume_event(msg, state) assert result["is_paused"] is False @@ -330,12 +526,9 @@ async def test_self_comment_is_ignored(self, worker): ) state = _prd_gate_state() - with patch("forge.orchestrator.worker.GitHubClient") as MockGH: - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh - + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) + with _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter): result = await worker._handle_resume_event(msg, state) # Should remain paused -- self-comment ignored @@ -358,15 +551,12 @@ async def test_self_comment_with_signature_is_ignored(self, worker): state = _prd_gate_state() settings = MagicMock(forge_bot_comment_prefix="my-signature") + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) with ( - patch("forge.orchestrator.worker.GitHubClient") as MockGH, + _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter), patch("forge.orchestrator.worker.get_settings", return_value=settings), ): - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh - result = await worker._handle_resume_event(msg, state) # Should remain paused -- self-comment with signature ignored @@ -389,15 +579,12 @@ async def test_own_comment_without_signature_is_not_ignored(self, worker): state = _prd_gate_state() settings = MagicMock(forge_bot_comment_prefix="my-signature") + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) with ( - patch("forge.orchestrator.worker.GitHubClient") as MockGH, + _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter), patch("forge.orchestrator.worker.get_settings", return_value=settings), ): - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh - result = await worker._handle_resume_event(msg, state) # Should be processed and no longer paused @@ -419,12 +606,9 @@ async def test_question_comment_sets_question_flag(self, worker): ) state = _prd_gate_state() - with patch("forge.orchestrator.worker.GitHubClient") as MockGH: - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh - + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) + with _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter): result = await worker._handle_resume_event(msg, state) assert result["is_paused"] is False @@ -441,6 +625,8 @@ async def test_inline_reply_resumes_only_matching_proposal_thread(self, worker): "comment": { "id": 12, "in_reply_to_id": 11, + "path": "prd.md", + "line": 20, "body": "Please make this change after all.", }, "sender": {"login": "reviewer"}, @@ -466,11 +652,9 @@ async def test_inline_reply_resumes_only_matching_proposal_thread(self, worker): ] ) - with patch("forge.orchestrator.worker.GitHubClient") as MockGH: - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) + with _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter): result = await worker._handle_resume_event(msg, state) assert result["revision_requested"] is True @@ -489,6 +673,8 @@ async def test_unknown_proposal_reply_target_is_ignored(self, worker, caplog): "comment": { "id": 31, "in_reply_to_id": 999, + "path": "prd.md", + "line": 5, "body": "This target is not in workflow state.", }, "sender": {"login": "reviewer"}, @@ -573,8 +759,10 @@ async def test_satisfied_bot_review_stays_paused(self, worker): ) state = _prd_gate_state(prd_content="# Current PRD") + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) with ( - patch("forge.orchestrator.worker.GitHubClient") as MockGH, + _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter), patch( "forge.orchestrator.worker.triage_automated_review", new=AsyncMock( @@ -584,10 +772,6 @@ async def test_satisfied_bot_review_stays_paused(self, worker): ), ) as triage, ): - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh result = await worker._handle_resume_event(msg, state) assert result == state @@ -606,8 +790,10 @@ async def test_blocking_bot_review_requests_bounded_revision(self, worker): ) state = _prd_gate_state(prd_content="# Current PRD") + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) with ( - patch("forge.orchestrator.worker.GitHubClient") as MockGH, + _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter), patch( "forge.orchestrator.worker.triage_automated_review", new=AsyncMock( @@ -619,10 +805,6 @@ async def test_blocking_bot_review_requests_bounded_revision(self, worker): ), ), ): - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh result = await worker._handle_resume_event(msg, state) assert result["revision_requested"] is True @@ -644,8 +826,10 @@ async def test_uncertain_bot_review_revises_with_original_feedback(self, worker) ) state = _prd_gate_state(prd_content="# Current PRD") + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) with ( - patch("forge.orchestrator.worker.GitHubClient") as MockGH, + _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter), patch( "forge.orchestrator.worker.triage_automated_review", new=AsyncMock( @@ -655,10 +839,6 @@ async def test_uncertain_bot_review_revises_with_original_feedback(self, worker) ), ), ): - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh result = await worker._handle_resume_event(msg, state) assert result["revision_requested"] is True @@ -679,8 +859,10 @@ async def test_bot_review_at_revision_cap_stays_paused(self, worker): ) state = _prd_gate_state(automated_review_revision_count=3) + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) with ( - patch("forge.orchestrator.worker.GitHubClient") as MockGH, + _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter), patch( "forge.orchestrator.worker.triage_automated_review", new=AsyncMock( @@ -690,10 +872,6 @@ async def test_bot_review_at_revision_cap_stays_paused(self, worker): ), ), ): - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh result = await worker._handle_resume_event(msg, state) assert result == state @@ -771,12 +949,9 @@ async def test_plain_comment_on_prd_pr_is_ignored(self, worker): ) state = _prd_gate_state() - with patch("forge.orchestrator.worker.GitHubClient") as MockGH: - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh - + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) + with _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter): result = await worker._handle_resume_event(msg, state) assert result.get("is_paused", True) is True @@ -799,12 +974,9 @@ async def test_bot_informational_comment_on_prd_pr_is_ignored(self, worker): ) state = _prd_gate_state() - with patch("forge.orchestrator.worker.GitHubClient") as MockGH: - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh - + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) + with _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter): result = await worker._handle_resume_event(msg, state) assert result.get("is_paused", True) is True @@ -825,29 +997,224 @@ async def test_human_review_bypasses_triage(self, worker): }, ) threads = [ - { - "thread_id": "thread-1", - "path": "prd.md", - "line": 10, - "comments": [{"comment_id": 100, "body": "Fix this section."}], - }, + Review( + id="thread-1", + state=ReviewState.COMMENTED, + body="", + author="", + comments=[ + ReviewComment( + id="100", path="prd.md", line=10, body="Fix this section.", author="" + ) + ], + ), ] state = _prd_gate_state(prd_content="# Current PRD") + mock_adapter = AsyncMock() + mock_adapter.get_review_thread_comments.return_value = threads + with ( - patch("forge.orchestrator.worker.GitHubClient") as MockGH, + _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter), patch( "forge.orchestrator.worker.triage_proposal_review_threads", new=AsyncMock(), ) as triage, ): - mock_gh = MagicMock() - mock_gh.get_pull_request_review_threads = AsyncMock(return_value=threads) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh - result = await worker._handle_resume_event(msg, state) assert result["revision_requested"] is True assert "Fix this section" in result["feedback_comment"] triage.assert_not_awaited() + + +class TestPrdPrReviewTypedFields: + """Brief Step 1: the PRD-PR review/merge branches read typed event fields.""" + + @pytest.mark.asyncio + async def test_review_with_changes_requested_sets_feedback(self, worker): + event = _make_normalized_event(kind=EventKind.REVIEW_SUBMITTED) + event.review = Review( + id="1", + state=ReviewState.CHANGES_REQUESTED, + body="please fix X", + author="reviewer1", + ) + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "prd_approval_gate", + "is_paused": True, + "prd_pr_number": 42, + "prd_pr_repo": "acme/payments", + } + + mock_adapter = AsyncMock() + mock_adapter.get_review_thread_comments.return_value = [] + with _patch_adapter(_repo_ref_for("acme/payments"), mock_adapter): + updated = await worker._handle_resume_event(message, current_state) + + assert updated["revision_requested"] is True + assert "please fix X" in updated["feedback_comment"] + + @pytest.mark.asyncio + async def test_pr_merged_sets_approved(self, worker): + event = _make_normalized_event(kind=EventKind.CR_MERGED) + event.change_request.state = ChangeRequestState.MERGED + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="cr_merged", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "prd_approval_gate", + "is_paused": True, + "prd_pr_number": 42, + "prd_pr_repo": "acme/payments", + } + + with patch("forge.orchestrator.worker.JiraClient") as MockJira: + MockJira.return_value.set_workflow_label = AsyncMock() + MockJira.return_value.close = AsyncMock() + updated = await worker._handle_resume_event(message, current_state) + + assert updated["is_paused"] is False + + +class TestProposalReplyTypedFields: + """The inline proposal-reply block reads typed ReviewComment fields. + + An inline pull_request_review_comment is distinguished from an issue comment + by comment.path being set; sender identity comes from actor.login. + """ + + def _reply_message(self, comment: ReviewComment, actor_login: str = "reviewer") -> QueueMessage: + event = _make_normalized_event(kind=EventKind.COMMENT_CREATED) + event.actor = Actor(login=actor_login, is_bot="[bot]" in actor_login) + event.comment = comment + return QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="comment_created", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + + def _state(self, **overrides) -> dict: + base = { + "current_node": "prd_approval_gate", + "is_paused": True, + "prd_pr_number": 42, + "prd_pr_repo": "acme/payments", + } + base.update(overrides) + return base + + @pytest.mark.asyncio + async def test_reply_matching_stored_decision_updates_and_unpauses(self, worker): + message = self._reply_message( + ReviewComment( + id="12", + body="Please make this change after all.", + author="reviewer", + path="prd.md", + line=20, + in_reply_to="11", + ) + ) + state = self._state( + proposal_review_decisions=[ + { + "thread_id": "thread-a", + "comment_id": 10, + "forge_reply_id": 11, + "disposition": "reply", + "feedback": "", + "response": "This conflicts with the API.", + }, + { + "thread_id": "thread-b", + "comment_id": 20, + "disposition": "reply", + "feedback": "", + "response": "Out of scope.", + }, + ] + ) + + with patch.object( + worker, "_get_forge_github_login", new=AsyncMock(return_value="forge-bot") + ): + result = await worker._handle_resume_event(message, state) + + assert result["is_paused"] is False + assert result["revision_requested"] is True + assert result["feedback_comment"] == "Please make this change after all." + assert result["proposal_review_decisions"][0]["disposition"] == "accept" + assert result["proposal_review_decisions"][0]["comment_id"] == 12 + assert result["proposal_review_decisions"][1] == state["proposal_review_decisions"][1] + + @pytest.mark.asyncio + async def test_standalone_reply_builds_thread_and_sets_rejection(self, worker): + # A human (non-bot) inline comment with no in_reply_to builds a fresh + # proposal_review_threads entry and requests a revision; the bot-only + # triage blocks are skipped for a human sender. + message = self._reply_message( + ReviewComment( + id="30", + body="Clarify the authorization behavior.", + author="reviewer", + path="prd.md", + line=12, + ), + actor_login="reviewer", + ) + state = self._state(prd_content="# Current PRD") + + with patch.object( + worker, "_get_forge_github_login", new=AsyncMock(return_value="forge-bot") + ): + result = await worker._handle_resume_event(message, state) + + assert result["is_paused"] is False + assert result["revision_requested"] is True + assert result["feedback_comment"] == "Clarify the authorization behavior." + + @pytest.mark.asyncio + async def test_self_reply_is_ignored(self, worker): + message = self._reply_message( + ReviewComment( + id="99", + body="Addressed in the latest revision.", + author="forge-bot", + path="prd.md", + line=3, + in_reply_to="11", + ), + actor_login="forge-bot", + ) + state = self._state( + proposal_review_decisions=[ + {"thread_id": "thread-a", "comment_id": 10, "forge_reply_id": 11} + ] + ) + + with patch.object( + worker, "_get_forge_github_login", new=AsyncMock(return_value="forge-bot") + ): + result = await worker._handle_resume_event(message, state) + + assert result == state diff --git a/tests/unit/orchestrator/test_worker_spec_pr.py b/tests/unit/orchestrator/test_worker_spec_pr.py index b695ff4e7..44deabb3f 100644 --- a/tests/unit/orchestrator/test_worker_spec_pr.py +++ b/tests/unit/orchestrator/test_worker_spec_pr.py @@ -1,24 +1,211 @@ """Tests for spec PR event handling in the worker.""" +from datetime import UTC, datetime from unittest.mock import AsyncMock, MagicMock, patch import pytest +from forge.integrations.source_control.contracts import ( + Actor, + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + EventKind, + NormalizedEvent, + Provider, + RepositoryRef, + Review, + ReviewComment, + ReviewState, +) from forge.models.events import EventSource from forge.orchestrator.worker import OrchestratorWorker -from forge.queue.models import QueueMessage +from forge.queue.models import QueueMessage, normalized_event_to_dict from forge.workflow.utils.automated_review_triage import AutomatedReviewDecision +def _repo_ref_for(namespace: str) -> RepositoryRef: + return RepositoryRef( + id=namespace, + provider=Provider.GITHUB, + connection="default-github", + namespace=namespace, + default_branch="main", + change_request_mode="fork", + ) + + +def _patch_adapter(repo_ref: RepositoryRef, adapter): + """Patch worker.get_adapter to resolve to the given (repo_ref, adapter) pair.""" + return patch("forge.orchestrator.worker.get_adapter", return_value=(repo_ref, adapter)) + + +_REVIEW_STATES = { + "approved": ReviewState.APPROVED, + "changes_requested": ReviewState.CHANGES_REQUESTED, + "commented": ReviewState.COMMENTED, + "pending": ReviewState.PENDING, +} + + +def _normalized_from_payload(event_type: str, payload: dict) -> NormalizedEvent: + """Build the NormalizedEvent a GitHub webhook payload would produce. + + Mirrors GitHubAdapter.parse_webhook for the event types exercised here so the + typed detection in _handle_resume_event runs against realistic data while the + raw payload is still carried for the triage blocks that read it. + """ + base_type = event_type.split(":", 1)[0] + repo_full = payload.get("repository", {}).get("full_name", "") + sender = payload.get("sender", {}) + sender_login = sender.get("login", "") + actor = Actor( + login=sender_login, + is_bot=sender.get("type") == "Bot" or "[bot]" in sender_login, + ) + repo_ref = RepositoryRef( + id=repo_full, + provider=Provider.GITHUB, + connection="default-github", + namespace=repo_full, + default_branch="main", + change_request_mode="fork", + ) + + change_request = None + pr = payload.get("pull_request") + issue = payload.get("issue") + if pr is not None: + if pr.get("merged", False): + state = ChangeRequestState.MERGED + elif pr.get("state") == "closed": + state = ChangeRequestState.CLOSED + else: + state = ChangeRequestState.OPEN + change_request = ChangeRequest( + identity=ChangeRequestIdentity( + connection="default-github", + repository_id=repo_full, + native_id=pr.get("number"), + ), + url=pr.get("html_url", ""), + title=pr.get("title", ""), + body=pr.get("body", "") or "", + state=state, + source_branch="", + target_branch="", + draft=False, + ) + elif issue is not None: + change_request = ChangeRequest( + identity=ChangeRequestIdentity( + connection="default-github", + repository_id=repo_full, + native_id=issue.get("number"), + ), + url=issue.get("html_url", ""), + title=issue.get("title", ""), + body=issue.get("body", "") or "", + state=ChangeRequestState.OPEN, + source_branch="", + target_branch="", + draft=False, + ) + + comment = None + review = None + if base_type in ("issue_comment", "pull_request_review_comment"): + kind = EventKind.COMMENT_CREATED + raw_comment = payload.get("comment", {}) + in_reply = raw_comment.get("in_reply_to_id") + comment = ReviewComment( + id=str(raw_comment.get("id", "")), + body=raw_comment.get("body", "") or "", + author=(raw_comment.get("user") or {}).get("login", ""), + path=raw_comment.get("path"), + line=raw_comment.get("line"), + in_reply_to=str(in_reply) if in_reply is not None else None, + ) + elif base_type == "pull_request_review": + kind = EventKind.REVIEW_SUBMITTED + raw_review = payload.get("review", {}) + review = Review( + id=str(raw_review.get("id", "")), + state=_REVIEW_STATES.get( + (raw_review.get("state") or "").lower(), ReviewState.COMMENTED + ), + body=raw_review.get("body", "") or "", + author=(raw_review.get("user") or {}).get("login", ""), + comments=[], + ) + elif base_type == "pull_request": + if change_request is not None and change_request.state == ChangeRequestState.MERGED: + kind = EventKind.CR_MERGED + elif change_request is not None and change_request.state == ChangeRequestState.CLOSED: + kind = EventKind.CR_CLOSED + else: + kind = EventKind.CR_UPDATED + else: + kind = EventKind.CR_UPDATED + + return NormalizedEvent( + id="evt-1", + kind=kind, + repo_ref=repo_ref, + actor=actor, + received_at=datetime(2026, 1, 1, tzinfo=UTC), + change_request=change_request, + comment=comment, + review=review, + raw=payload, + ) + + def _make_message(event_type: str, payload: dict, ticket_key: str = "TEST-123") -> QueueMessage: return QueueMessage( message_id="msg-1", event_id="evt-1", - source=EventSource.GITHUB, + source=EventSource.SOURCE_CONTROL, event_type=event_type, ticket_key=ticket_key, payload=payload, + normalized_event=normalized_event_to_dict(_normalized_from_payload(event_type, payload)), + ) + + +def _make_normalized_event(**overrides) -> NormalizedEvent: + """A canned NormalizedEvent (repo acme/payments, PR 42) for typed-field tests.""" + repo_ref = RepositoryRef( + id="acme/payments", + provider=Provider.GITHUB, + connection="default-github", + namespace="acme/payments", + default_branch="main", + change_request_mode="fork", + ) + change_request = ChangeRequest( + identity=ChangeRequestIdentity( + connection="default-github", repository_id="acme/payments", native_id=42 + ), + url="https://github.com/acme/payments/pull/42", + title="t", + body="", + state=ChangeRequestState.OPEN, + source_branch="feature", + target_branch="main", + draft=False, ) + defaults = { + "id": "delivery-1", + "kind": EventKind.CR_OPENED, + "repo_ref": repo_ref, + "actor": Actor(login="octocat", is_bot=False), + "received_at": datetime(2026, 1, 1, tzinfo=UTC), + "change_request": change_request, + "raw": {}, + } + defaults.update(overrides) + return NormalizedEvent(**defaults) def _spec_gate_state(**overrides) -> dict: @@ -49,6 +236,7 @@ def worker(): w = OrchestratorWorker.__new__(OrchestratorWorker) w._post_terminal_error_comment = AsyncMock() w._post_resume_ack_comment = AsyncMock() + w._forge_github_logins = {} return w @@ -107,18 +295,80 @@ async def test_satisfied_bot_spec_review_stays_paused(worker): ) state = _spec_gate_state() + mock_adapter = AsyncMock() + mock_adapter.get_authenticated_identity.return_value = Actor(login="forge-bot", is_bot=True) with ( - patch("forge.orchestrator.worker.GitHubClient") as MockGH, + _patch_adapter(_repo_ref_for("org/proposals"), mock_adapter), patch( "forge.orchestrator.worker.triage_automated_review", new=AsyncMock(return_value=AutomatedReviewDecision("satisfied")), ) as triage, ): - mock_gh = MagicMock() - mock_gh.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - mock_gh.close = AsyncMock() - MockGH.return_value = mock_gh result = await worker._handle_resume_event(msg, state) assert result == state triage.assert_awaited_once() + + +class TestSpecPrReviewTypedFields: + """Brief Step 1 (spec mirror): spec-PR review/merge read typed event fields.""" + + @pytest.mark.asyncio + async def test_review_with_changes_requested_sets_feedback(self, worker): + event = _make_normalized_event(kind=EventKind.REVIEW_SUBMITTED) + event.review = Review( + id="1", + state=ReviewState.CHANGES_REQUESTED, + body="please fix X", + author="reviewer1", + ) + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="review_submitted", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "spec_approval_gate", + "is_paused": True, + "spec_pr_number": 42, + "spec_pr_repo": "acme/payments", + } + + mock_adapter = AsyncMock() + mock_adapter.get_review_thread_comments.return_value = [] + with _patch_adapter(_repo_ref_for("acme/payments"), mock_adapter): + updated = await worker._handle_resume_event(message, current_state) + + assert updated["revision_requested"] is True + assert "please fix X" in updated["feedback_comment"] + + @pytest.mark.asyncio + async def test_pr_merged_sets_approved(self, worker): + event = _make_normalized_event(kind=EventKind.CR_MERGED) + event.change_request.state = ChangeRequestState.MERGED + message = QueueMessage( + message_id="1", + event_id="e1", + source=EventSource.SOURCE_CONTROL, + event_type="cr_merged", + ticket_key="PROJ-1", + payload={}, + normalized_event=normalized_event_to_dict(event), + ) + current_state = { + "current_node": "spec_approval_gate", + "is_paused": True, + "spec_pr_number": 42, + "spec_pr_repo": "acme/payments", + } + + with patch("forge.orchestrator.worker.JiraClient") as MockJira: + MockJira.return_value.set_workflow_label = AsyncMock() + MockJira.return_value.close = AsyncMock() + updated = await worker._handle_resume_event(message, current_state) + + assert updated["is_paused"] is False diff --git a/tests/unit/queue/test_consumer.py b/tests/unit/queue/test_consumer.py index 7ba21fcde..c0ed1fba1 100644 --- a/tests/unit/queue/test_consumer.py +++ b/tests/unit/queue/test_consumer.py @@ -490,3 +490,44 @@ async def handler(message: QueueMessage) -> None: "TICKET-B (fast) should have completed before TICKET-A (slow). " "This is the AISOS-709 regression — sequential processing detected." ) + + +class TestLegacyStreamMigration: + """The pre-rename source-control stream is never auto-consumed. + + Its entries predate the NormalizedEvent/adapter cutover and have no + normalized_event to deserialize, so processing them would silently ack + away whatever CI/review/merge signal they carried instead of acting on + it. They're left for a deliberate, out-of-band migration; only + forge:events:source_control (and forge:events:jira) are consumed live. + """ + + @pytest.mark.asyncio + async def test_ensure_consumer_groups_does_not_include_legacy_stream(self) -> None: + redis_mock = _make_redis_mock() + consumer = _make_consumer(redis_mock) + + await consumer._ensure_consumer_groups() + + created_streams = {call.args[0] for call in redis_mock.xgroup_create.await_args_list} + assert "forge:events:github" not in created_streams + assert "forge:events:source_control" in created_streams + + @pytest.mark.asyncio + async def test_start_does_not_consume_legacy_stream(self) -> None: + redis_mock = _make_redis_mock() + consumer = _make_consumer(redis_mock) + consumer.register_handler(EventSource.SOURCE_CONTROL, AsyncMock()) + + consumed_streams: list[str] = [] + + async def fake_consume_stream(stream: str, _source: EventSource) -> None: + consumed_streams.append(stream) + + consumer._consume_stream = fake_consume_stream + consumer._process_retry_queue = AsyncMock() + + await consumer.start() + + assert "forge:events:source_control" in consumed_streams + assert "forge:events:github" not in consumed_streams diff --git a/tests/unit/queue/test_normalized_event_transport.py b/tests/unit/queue/test_normalized_event_transport.py new file mode 100644 index 000000000..564b850eb --- /dev/null +++ b/tests/unit/queue/test_normalized_event_transport.py @@ -0,0 +1,167 @@ +"""Round-trip serialization tests for NormalizedEvent <-> QueueMessage.""" + +from datetime import UTC, datetime + +from forge.integrations.source_control.contracts import ( + Actor, + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + CheckConclusion, + CheckRun, + CheckStatus, + EventKind, + NormalizedEvent, + Provider, + RepositoryRef, + Review, + ReviewComment, + ReviewState, +) +from forge.models.events import EventSource +from forge.queue.models import QueueMessage, normalized_event_from_dict, normalized_event_to_dict + + +def _sample_event() -> NormalizedEvent: + repo_ref = RepositoryRef( + id="test/repo", + provider=Provider.GITHUB, + connection="default-github", + namespace="test/repo", + default_branch="main", + change_request_mode="fork", + ) + change_request = ChangeRequest( + identity=ChangeRequestIdentity( + connection="default-github", repository_id="test/repo", native_id=42 + ), + url="https://github.com/test/repo/pull/42", + title="Test PR", + body="body", + state=ChangeRequestState.OPEN, + source_branch="feature", + target_branch="main", + head_sha="sha123", + draft=False, + ) + return NormalizedEvent( + id="delivery-123", + kind=EventKind.CR_OPENED, + repo_ref=repo_ref, + actor=Actor(login="octocat", is_bot=False), + received_at=datetime(2026, 1, 1, tzinfo=UTC), + change_request=change_request, + raw={"action": "opened"}, + ) + + +def test_normalized_event_round_trips_through_dict(): + event = _sample_event() + + data = normalized_event_to_dict(event) + restored = normalized_event_from_dict(data) + + assert restored.id == event.id + assert restored.kind == event.kind + assert restored.repo_ref == event.repo_ref + assert restored.actor == event.actor + assert restored.received_at == event.received_at + assert restored.change_request == event.change_request + assert restored.change_request.head_sha == "sha123" + assert restored.comment is None + assert restored.review is None + assert restored.check is None + assert restored.raw == event.raw + + +def test_queue_message_from_redis_maps_legacy_github_source(): + """Retry/DLQ entries and unconsumed stream messages persisted before the + EventSource "github" -> "source_control" rename must still decode.""" + message = QueueMessage.from_redis( + "1-0", + { + "event_id": "evt-1", + "source": "github", + "event_type": "cr_opened", + "ticket_key": "PROJ-1", + "payload": "{}", + "normalized_event": "", + "timestamp": "2026-01-01T00:00:00", + "retry_count": "0", + }, + ) + + assert message.source == EventSource.SOURCE_CONTROL + + +def test_normalized_event_round_trips_with_no_change_request(): + event = _sample_event() + event.change_request = None + + data = normalized_event_to_dict(event) + restored = normalized_event_from_dict(data) + + assert restored.change_request is None + + +def test_normalized_event_round_trips_populated_comment(): + event = _sample_event() + event.kind = EventKind.COMMENT_CREATED + event.comment = ReviewComment( + id="c1", + body="/forge skip-gate flaky-test", + author="octocat", + path="src/app.py", + line=12, + resolved=True, + in_reply_to="c0", + ) + + data = normalized_event_to_dict(event) + restored = normalized_event_from_dict(data) + + assert restored.comment == event.comment + + +def test_normalized_event_round_trips_populated_review(): + event = _sample_event() + event.kind = EventKind.REVIEW_SUBMITTED + event.review = Review( + id="r1", + state=ReviewState.CHANGES_REQUESTED, + body="please fix X", + author="reviewer1", + comments=[ + ReviewComment( + id="rc1", + body="inline nit", + author="reviewer1", + path="src/app.py", + line=7, + ) + ], + ) + + data = normalized_event_to_dict(event) + restored = normalized_event_from_dict(data) + + assert restored.review == event.review + + +def test_normalized_event_round_trips_check_output(): + event = _sample_event() + event.kind = EventKind.CHECK_UPDATED + event.check = CheckRun( + name="build", + status=CheckStatus.COMPLETED, + conclusion=CheckConclusion.FAILURE, + url="https://github.com/test/repo/runs/1", + logs_url="https://github.com/test/repo/runs/1/logs", + output={"title": "Build failed", "summary": "2 tests failed", "text": "details"}, + ) + + data = normalized_event_to_dict(event) + restored = normalized_event_from_dict(data) + + assert restored.check == event.check + assert restored.check.output == event.check.output diff --git a/tests/unit/queue/test_producer.py b/tests/unit/queue/test_producer.py index cf4374f36..56613b71a 100644 --- a/tests/unit/queue/test_producer.py +++ b/tests/unit/queue/test_producer.py @@ -1,10 +1,37 @@ """Unit tests for atomic queue publication and deduplication.""" +from datetime import UTC, datetime from unittest.mock import AsyncMock +from forge.integrations.source_control.contracts import ( + Actor, + EventKind, + NormalizedEvent, + Provider, + RepositoryRef, +) from forge.models.events import EventSource from forge.queue.deduplication import DEDUP_KEY_PREFIX, DEDUP_TTL_SECONDS -from forge.queue.producer import JIRA_STREAM, QueueProducer +from forge.queue.producer import JIRA_STREAM, SOURCE_CONTROL_STREAM, QueueProducer + + +def _sample_event() -> NormalizedEvent: + repo_ref = RepositoryRef( + id="acme/payments", + provider=Provider.GITHUB, + connection="default-github", + namespace="acme/payments", + default_branch="main", + change_request_mode="fork", + ) + return NormalizedEvent( + id="delivery-1", + kind=EventKind.CR_OPENED, + repo_ref=repo_ref, + actor=Actor(login="octocat", is_bot=False), + received_at=datetime(2026, 1, 1, tzinfo=UTC), + raw={"action": "opened"}, + ) async def test_publish_once_returns_stream_id_for_new_event() -> None: @@ -44,3 +71,33 @@ async def test_publish_once_returns_none_for_duplicate_event() -> None: ) assert message_id is None + + +async def test_publish_event_returns_stream_id_for_new_event() -> None: + """publish_event must use the same atomic dedup as publish_once -- GitHub + redelivers webhooks on timeouts/5xx/manual redeliver, and a plain XADD + would silently reprocess every retried delivery.""" + redis_client = AsyncMock() + redis_client.eval.return_value = "456-0" + producer = QueueProducer(redis_client=redis_client) + + message_id = await producer.publish_event(_sample_event(), ticket_key="PROJ-1") + + assert message_id == "456-0" + args = redis_client.eval.await_args.args + assert args[1:5] == ( + 2, + f"{DEDUP_KEY_PREFIX}delivery-1", + SOURCE_CONTROL_STREAM, + DEDUP_TTL_SECONDS, + ) + + +async def test_publish_event_returns_none_for_duplicate_event() -> None: + redis_client = AsyncMock() + redis_client.eval.return_value = None + producer = QueueProducer(redis_client=redis_client) + + message_id = await producer.publish_event(_sample_event(), ticket_key="PROJ-1") + + assert message_id is None diff --git a/tests/unit/workflow/nodes/test_ci_attempt_tracking.py b/tests/unit/workflow/nodes/test_ci_attempt_tracking.py index 9cad5d96e..c9e412d6d 100644 --- a/tests/unit/workflow/nodes/test_ci_attempt_tracking.py +++ b/tests/unit/workflow/nodes/test_ci_attempt_tracking.py @@ -1,23 +1,66 @@ """Unit tests for CI attempt tracking (AISOS-654).""" -import pytest from unittest.mock import AsyncMock, MagicMock, patch +import pytest + +from forge.integrations.source_control.contracts import ( + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + CheckConclusion, + CheckRun, + CheckStatus, + Provider, + RepositoryRef, +) from forge.models.workflow import ForgeLabel -from forge.workflow.nodes.ci_evaluator import evaluate_ci_status from forge.workflow.feature.state import FeatureState - +from forge.workflow.nodes.ci_evaluator import evaluate_ci_status # ── Helpers ─────────────────────────────────────────────────────────────────── -def create_mock_github_client(): - """Create a mock GitHub client with common methods.""" - client = MagicMock() - client.get_pull_request = AsyncMock() - client.get_check_runs = AsyncMock() - client.close = AsyncMock() - return client +def _repo_ref(repo: str = "org/repo") -> RepositoryRef: + return RepositoryRef( + id=repo, + provider=Provider.GITHUB, + connection="c", + namespace=repo, + default_branch="main", + change_request_mode="fork", + ) + + +def _mock_adapter(status: str, conclusion: str | None, pr_number: int = 42, repo: str = "org/repo"): + """Patch target for ci_evaluator.get_adapter returning one check with the given result.""" + adapter = AsyncMock() + adapter.get_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity(connection="c", repository_id=repo, native_id=pr_number), + url=f"https://github.com/{repo}/pull/{pr_number}", + title="t", + body="", + state=ChangeRequestState.OPEN, + source_branch="feature", + target_branch="main", + head_sha="deadbeef", + ) + ) + adapter.get_checks = AsyncMock( + return_value=[ + CheckRun( + name="test", + status=CheckStatus.COMPLETED if status == "completed" else CheckStatus.IN_PROGRESS, + conclusion={ + "failure": CheckConclusion.FAILURE, + "success": CheckConclusion.SUCCESS, + None: CheckConclusion.NONE, + }[conclusion], + ) + ] + ) + return adapter def create_base_state(**kwargs) -> FeatureState: @@ -31,11 +74,32 @@ def create_base_state(**kwargs) -> FeatureState: "ci_failed_checks": [], "ci_skipped_checks": [], "current_repo": "org/repo", + "current_pr_number": 42, } defaults.update(kwargs) return FeatureState(**defaults) +@pytest.mark.asyncio +async def test_evaluate_ci_status_queries_checks_by_head_sha_not_branch() -> None: + """CI checks must be looked up by the PR head SHA, not its source branch name. + + In fork mode the head branch only exists on the fork, so GitHub cannot + resolve it against the upstream repository the checks are queried against. + """ + adapter = _mock_adapter("completed", "failure") + state = create_base_state() + + with ( + patch("forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(_repo_ref(), adapter)), + patch("forge.integrations.jira.client.JiraClient.close", new=AsyncMock()), + patch("forge.workflow.utils.jira_status.set_review_pending_label", new=AsyncMock()), + ): + await evaluate_ci_status(state) + + adapter.get_checks.assert_awaited_once_with(_repo_ref(), "deadbeef") + + @pytest.mark.asyncio async def test_multi_pr_ci_rejects_inconsistent_active_view() -> None: state = create_base_state( @@ -120,23 +184,17 @@ 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 = [ - { - "name": "test", - "status": "completed", - "conclusion": "failure", - "output": {}, - "html_url": "https://github.com/org/repo/runs/1", - } - ] + adapter = _mock_adapter("completed", "failure") - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=github): - 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"] - result = await evaluate_ci_status(state) + with ( + patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(_repo_ref(), adapter) + ), + 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"] + result = await evaluate_ci_status(state) assert result["ci_fix_attempt"] == 1 assert result["current_node"] == "attempt_ci_fix" @@ -146,23 +204,17 @@ 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 = [ - { - "name": "test", - "status": "completed", - "conclusion": "failure", - "output": {}, - "html_url": "https://github.com/org/repo/runs/1", - } - ] + adapter = _mock_adapter("completed", "failure") - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=github): - 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"] - result = await evaluate_ci_status(state) + with ( + patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(_repo_ref(), adapter) + ), + 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"] + result = await evaluate_ci_status(state) assert result["ci_fix_attempt"] == 2 assert result["current_node"] == "attempt_ci_fix" @@ -172,23 +224,17 @@ 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 = [ - { - "name": "test", - "status": "completed", - "conclusion": "failure", - "output": {}, - "html_url": "https://github.com/org/repo/runs/1", - } - ] + adapter = _mock_adapter("completed", "failure") - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=github): - 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"] - result = await evaluate_ci_status(state) + with ( + patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(_repo_ref(), adapter) + ), + 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"] + result = await evaluate_ci_status(state) assert result["ci_fix_attempt"] == 3 assert result["current_node"] == "attempt_ci_fix" @@ -205,26 +251,18 @@ 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 = [ - { - "name": "test", - "status": "completed", - "conclusion": "failure", - "output": {}, - "html_url": "https://github.com/org/repo/runs/1", - } - ] + adapter = _mock_adapter("completed", "failure") - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=github): - 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: - result = await evaluate_ci_status(state) + with ( + patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(_repo_ref(), adapter) + ), + 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: + result = await evaluate_ci_status(state) # Should not increment or route to attempt_ci_fix assert result["ci_fix_attempt"] == 3 # Unchanged @@ -237,26 +275,18 @@ 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 = [ - { - "name": "test", - "status": "completed", - "conclusion": "failure", - "output": {}, - "html_url": "https://github.com/org/repo/runs/1", - } - ] + adapter = _mock_adapter("completed", "failure") - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=github): - 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: - result = await evaluate_ci_status(state) + with ( + patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(_repo_ref(), adapter) + ), + 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: + result = await evaluate_ci_status(state) # Should not increment or route to attempt_ci_fix assert result["ci_fix_attempt"] == 4 # Unchanged @@ -269,23 +299,17 @@ 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 = [ - { - "name": "test", - "status": "completed", - "conclusion": "failure", - "output": {}, - "html_url": "https://github.com/org/repo/runs/1", - } - ] + adapter = _mock_adapter("completed", "failure") - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=github): - 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"] - result = await evaluate_ci_status(state) + with ( + patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(_repo_ref(), adapter) + ), + 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"] + result = await evaluate_ci_status(state) # Should increment and route to attempt_ci_fix assert result["ci_fix_attempt"] == 3 @@ -304,24 +328,16 @@ 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 = [ - { - "name": "test", - "status": "completed", - "conclusion": "success", - "output": {}, - "html_url": "https://github.com/org/repo/runs/1", - } - ] + adapter = _mock_adapter("completed", "success") jira = MagicMock() jira.set_workflow_label = AsyncMock() jira.close = AsyncMock() with ( - patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=github), + patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(_repo_ref(), adapter) + ), patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=jira), patch("forge.workflow.nodes.ci_evaluator.get_settings") as mock_settings, ): @@ -372,23 +388,17 @@ async def test_missing_current_attempt_defaults_to_zero(self): # 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 = [ - { - "name": "test", - "status": "completed", - "conclusion": "failure", - "output": {}, - "html_url": "https://github.com/org/repo/runs/1", - } - ] + adapter = _mock_adapter("completed", "failure") - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=github): - 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"] - result = await evaluate_ci_status(state) + with ( + patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(_repo_ref(), adapter) + ), + 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"] + result = await evaluate_ci_status(state) # Should default to 0 and increment to 1 assert result["ci_fix_attempt"] == 1 @@ -400,23 +410,17 @@ async def test_missing_max_attempts_defaults_to_config_value(self): # 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 = [ - { - "name": "test", - "status": "completed", - "conclusion": "failure", - "output": {}, - "html_url": "https://github.com/org/repo/runs/1", - } - ] + adapter = _mock_adapter("completed", "failure") - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=github): - 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"] - result = await evaluate_ci_status(state) + with ( + patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(_repo_ref(), adapter) + ), + 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"] + result = await evaluate_ci_status(state) # Should allow attempt since default is 5 assert result["ci_fix_attempt"] == 1 @@ -427,23 +431,17 @@ 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 = [ - { - "name": "test", - "status": "completed", - "conclusion": "failure", - "output": {}, - "html_url": "https://github.com/org/repo/runs/1", - } - ] + adapter = _mock_adapter("completed", "failure") - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=github): - 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"] - result = await evaluate_ci_status(state) + with ( + patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(_repo_ref(), adapter) + ), + 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"] + result = await evaluate_ci_status(state) # Should allow first attempt assert result["ci_fix_attempt"] == 1 @@ -451,12 +449,16 @@ async def test_max_attempts_one_allows_single_attempt(self): # Second failure should block state2 = create_base_state(ci_fix_attempt=1, ci_fix_max_attempts=1) - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=github): - 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"): - result2 = await evaluate_ci_status(state2) + with ( + patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(_repo_ref(), adapter) + ), + 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"): + result2 = await evaluate_ci_status(state2) assert result2["ci_fix_attempt"] == 1 # Unchanged assert result2["current_node"] == "ci_evaluator" diff --git a/tests/unit/workflow/nodes/test_ci_attribution.py b/tests/unit/workflow/nodes/test_ci_attribution.py index a01df0d38..bf27fbe3d 100644 --- a/tests/unit/workflow/nodes/test_ci_attribution.py +++ b/tests/unit/workflow/nodes/test_ci_attribution.py @@ -1,11 +1,54 @@ """Tests for CI attribution and evaluate_ci_status pending_ci_event clearing.""" -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, patch import pytest +from forge.integrations.source_control.contracts import ( + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + CheckConclusion, + CheckRun, + CheckStatus, + Provider, + RepositoryRef, +) + + +def _repo_ref(repo: str = "org/repo") -> RepositoryRef: + return RepositoryRef( + id=repo, + provider=Provider.GITHUB, + connection="c", + namespace=repo, + default_branch="main", + change_request_mode="fork", + ) + + +def _mock_adapter(checks, pr_number=1, repo="org/repo"): + """Patch target for ci_evaluator.get_adapter returning the given checks.""" + adapter = AsyncMock() + adapter.get_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity(connection="c", repository_id=repo, native_id=pr_number), + url=f"https://github.com/{repo}/pull/{pr_number}", + title="t", + body="", + state=ChangeRequestState.OPEN, + source_branch="feature", + target_branch="main", + ) + ) + adapter.get_checks = AsyncMock(return_value=checks) + return adapter + + BASE_STATE = { "ticket_key": "TEST-1", + "current_repo": "org/repo", + "current_pr_number": 1, "pr_urls": ["https://github.com/org/repo/pull/1"], "ci_fix_attempt": 0, "ci_fix_max_attempts": 5, @@ -28,16 +71,12 @@ @pytest.mark.asyncio -@patch("forge.workflow.nodes.ci_evaluator.GitHubClient") -async def test_evaluate_ci_status_clears_pending_ci_event_on_pending(MockGitHub): +@patch("forge.workflow.nodes.ci_evaluator.get_adapter") +async def test_evaluate_ci_status_clears_pending_ci_event_on_pending(mock_get_adapter): """evaluate_ci_status returns pending_ci_event=False when CI is still running.""" from forge.workflow.nodes.ci_evaluator import evaluate_ci_status - mock_github = AsyncMock() - MockGitHub.return_value = mock_github - mock_github.get_pull_request = AsyncMock(return_value={"head": {"sha": "abc123"}}) - mock_github.get_check_runs = AsyncMock(return_value=[]) # No checks yet = pending - mock_github.close = AsyncMock() + mock_get_adapter.return_value = (_repo_ref(), _mock_adapter([])) # No checks yet = pending state = {**BASE_STATE, "pending_ci_event": True} result = await evaluate_ci_status(state) @@ -47,19 +86,24 @@ async def test_evaluate_ci_status_clears_pending_ci_event_on_pending(MockGitHub) @pytest.mark.asyncio -@patch("forge.workflow.nodes.ci_evaluator.GitHubClient") +@patch("forge.workflow.nodes.ci_evaluator.get_adapter") @patch("forge.workflow.nodes.ci_evaluator.JiraClient") -async def test_evaluate_ci_status_clears_pending_ci_event_on_passed(MockJira, MockGitHub): +async def test_evaluate_ci_status_clears_pending_ci_event_on_passed(MockJira, mock_get_adapter): """evaluate_ci_status returns pending_ci_event=False when all CI passes.""" from forge.workflow.nodes.ci_evaluator import evaluate_ci_status - mock_github = AsyncMock() - MockGitHub.return_value = mock_github - mock_github.get_pull_request = AsyncMock(return_value={"head": {"sha": "abc123"}}) - mock_github.get_check_runs = AsyncMock( - return_value=[{"name": "unit-tests", "status": "completed", "conclusion": "success"}] + mock_get_adapter.return_value = ( + _repo_ref(), + _mock_adapter( + [ + CheckRun( + name="unit-tests", + status=CheckStatus.COMPLETED, + conclusion=CheckConclusion.SUCCESS, + ) + ] + ), ) - mock_github.close = AsyncMock() mock_jira = AsyncMock() MockJira.return_value = mock_jira @@ -72,14 +116,42 @@ async def test_evaluate_ci_status_clears_pending_ci_event_on_passed(MockJira, Mo assert result.get("pending_ci_event", True) is False +@pytest.mark.asyncio +@patch("forge.workflow.nodes.ci_evaluator.get_adapter") +async def test_evaluate_ci_status_finds_url_keyed_pr_when_number_unknown(mock_get_adapter): + """A PR saved before its number was known is stored under a URL key + (see pr_state.py). evaluate_ci_status must locate it via that fallback + instead of failing with "Active pull request state is inconsistent".""" + from forge.workflow.nodes.ci_evaluator import evaluate_ci_status + + mock_get_adapter.return_value = (_repo_ref(), _mock_adapter([])) # No checks yet = pending + + pr_url = "https://github.com/org/repo/pull/1" + state = { + **BASE_STATE, + "current_repo": "org/repo", + "current_pr_url": pr_url, + "current_pr_number": None, + "pull_requests": { + f"org/repo:{pr_url}": {"url": pr_url, "number": None, "repo": "org/repo"} + }, + "pending_ci_event": True, + } + + result = await evaluate_ci_status(state) + + assert result.get("ci_status") == "pending" + assert result.get("last_error") is None + + @pytest.mark.asyncio @patch("forge.workflow.nodes.ci_evaluator.ContainerRunner") @patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") @patch("forge.workflow.nodes.ci_evaluator.JiraClient") -@patch("forge.workflow.nodes.ci_evaluator.GitHubClient") +@patch("forge.workflow.nodes.ci_evaluator.get_adapter") @patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", new_callable=AsyncMock) async def test_attribution_external_skips_fix( - _mock_fetch, MockGitHub, MockJira, mock_prep, MockRunner, tmp_path + _mock_fetch, mock_get_adapter, MockJira, mock_prep, MockRunner, tmp_path ): """Phase 0 verdict attributable=false → external_failure, no fix attempt increment.""" from forge.workflow.nodes.ci_evaluator import attempt_ci_fix @@ -105,8 +177,7 @@ async def write_attribution(**_kwargs): MockJira.return_value = mock_jira mock_jira.close = AsyncMock() - mock_github = AsyncMock() - MockGitHub.return_value = mock_github + mock_get_adapter.return_value = (_repo_ref(), AsyncMock()) state = {**ATTEMPT_BASE_STATE, "workspace_path": str(tmp_path)} result = await attempt_ci_fix(state) @@ -135,10 +206,10 @@ def test_attribution_prompt_compares_full_pr_against_default_branch(): @patch("forge.workflow.nodes.ci_evaluator.ContainerRunner") @patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") @patch("forge.workflow.nodes.ci_evaluator.JiraClient") -@patch("forge.workflow.nodes.ci_evaluator.GitHubClient") +@patch("forge.workflow.nodes.ci_evaluator.get_adapter") @patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", new_callable=AsyncMock) async def test_attribution_attributable_proceeds_to_fix( - _mock_fetch, MockGitHub, MockJira, mock_prep, MockRunner, tmp_path + _mock_fetch, mock_get_adapter, MockJira, mock_prep, MockRunner, tmp_path ): """Phase 0 verdict attributable=true → proceeds to fix, increments attempt.""" from forge.workflow.nodes.ci_evaluator import attempt_ci_fix @@ -162,8 +233,7 @@ async def test_attribution_attributable_proceeds_to_fix( MockJira.return_value = mock_jira mock_jira.close = AsyncMock() - mock_github = AsyncMock() - MockGitHub.return_value = mock_github + mock_get_adapter.return_value = (_repo_ref(), AsyncMock()) state = {**ATTEMPT_BASE_STATE, "workspace_path": str(tmp_path)} result = await attempt_ci_fix(state) @@ -173,14 +243,30 @@ async def test_attribution_attributable_proceeds_to_fix( assert result.get("pending_ci_event", True) is False +@pytest.mark.asyncio +async def test_attempt_ci_fix_with_no_failed_checks_reverifies_instead_of_asserting_pass(): + """attempt_ci_fix should only ever be routed to with ci_failed_checks + populated. If it's unexpectedly empty (e.g. a concurrent/stale state + update cleared it), the node must re-verify live CI via ci_evaluator + rather than asserting ci_status=passed without checking.""" + from forge.workflow.nodes.ci_evaluator import attempt_ci_fix + + state = {**ATTEMPT_BASE_STATE, "ci_failed_checks": []} + result = await attempt_ci_fix(state) + + assert result["current_node"] == "ci_evaluator" + assert result.get("ci_status") != "passed" + assert result.get("pending_ci_event", True) is False + + @pytest.mark.asyncio @patch("forge.workflow.nodes.ci_evaluator.ContainerRunner") @patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") @patch("forge.workflow.nodes.ci_evaluator.JiraClient") -@patch("forge.workflow.nodes.ci_evaluator.GitHubClient") +@patch("forge.workflow.nodes.ci_evaluator.get_adapter") @patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", new_callable=AsyncMock) async def test_attribution_missing_file_proceeds_as_attributable( - _mock_fetch, MockGitHub, MockJira, mock_prep, MockRunner, tmp_path + _mock_fetch, mock_get_adapter, MockJira, mock_prep, MockRunner, tmp_path ): """Missing ci-attribution.json is treated as attributable (fail-safe).""" from forge.workflow.nodes.ci_evaluator import attempt_ci_fix @@ -196,8 +282,7 @@ async def test_attribution_missing_file_proceeds_as_attributable( mock_jira = AsyncMock() MockJira.return_value = mock_jira mock_jira.close = AsyncMock() - mock_github = AsyncMock() - MockGitHub.return_value = mock_github + mock_get_adapter.return_value = (_repo_ref(), AsyncMock()) state = {**ATTEMPT_BASE_STATE, "workspace_path": str(tmp_path)} result = await attempt_ci_fix(state) @@ -211,10 +296,10 @@ async def test_attribution_missing_file_proceeds_as_attributable( @patch("forge.workflow.nodes.ci_evaluator.ContainerRunner") @patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") @patch("forge.workflow.nodes.ci_evaluator.JiraClient") -@patch("forge.workflow.nodes.ci_evaluator.GitHubClient") +@patch("forge.workflow.nodes.ci_evaluator.get_adapter") @patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", new_callable=AsyncMock) async def test_attribution_does_not_reuse_stale_external_verdict( - _mock_fetch, MockGitHub, MockJira, mock_prep, MockRunner, tmp_path + _mock_fetch, mock_get_adapter, MockJira, mock_prep, MockRunner, tmp_path ): """A missing fresh verdict must not reuse an external verdict from an earlier CI cycle.""" from forge.workflow.nodes.ci_evaluator import attempt_ci_fix @@ -233,8 +318,7 @@ async def test_attribution_does_not_reuse_stale_external_verdict( mock_jira = AsyncMock() MockJira.return_value = mock_jira mock_jira.close = AsyncMock() - mock_github = AsyncMock() - MockGitHub.return_value = mock_github + mock_get_adapter.return_value = (_repo_ref(), AsyncMock()) result = await attempt_ci_fix({**ATTEMPT_BASE_STATE, "workspace_path": str(tmp_path)}) @@ -243,54 +327,6 @@ async def test_attribution_does_not_reuse_stale_external_verdict( assert not attribution_file.exists() -@pytest.mark.asyncio -@patch("forge.workflow.nodes.ci_evaluator.GitOperations") -@patch("forge.workflow.nodes.ci_evaluator.ContainerRunner") -@patch("forge.workflow.nodes.ci_evaluator.prepare_workspace") -@patch("forge.workflow.nodes.ci_evaluator.JiraClient") -@patch("forge.workflow.nodes.ci_evaluator.GitHubClient") -@patch("forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", new_callable=AsyncMock) -async def test_completed_ci_fix_captures_attempt_handoff( - _mock_fetch, MockGitHub, MockJira, mock_prep, MockRunner, MockGitOperations, tmp_path -): - """A completed CI fix checkpoints its handoff with the current attempt number.""" - from forge.workflow.nodes.ci_evaluator import attempt_ci_fix - - forge_dir = tmp_path / ".forge" - forge_dir.mkdir() - mock_prep.return_value = (str(tmp_path), None) - - run_count = 0 - - async def run_phase(**_kwargs): - nonlocal run_count - run_count += 1 - if run_count == 1: - (forge_dir / "ci-attribution.json").write_text('{"attributable": true}') - elif run_count == 2: - (forge_dir / "fix-plan.md").write_text("apply the fix") - else: - (forge_dir / "handoff.md").write_text("CI fix completed") - return MagicMock(success=True) - - MockRunner.return_value.run = AsyncMock(side_effect=run_phase) - MockJira.return_value = AsyncMock() - MockGitHub.return_value = AsyncMock() - - mock_git = MagicMock() - mock_git.has_uncommitted_changes.return_value = False - mock_git._run_git.return_value.stdout = "" - MockGitOperations.return_value = mock_git - - result = await attempt_ci_fix({**ATTEMPT_BASE_STATE, "workspace_path": str(tmp_path)}) - - assert result["current_node"] == "human_review_gate" - assert result["last_error"] is None - handoff = result["handoffs"]["org/repo"] - assert handoff["content"] == "CI fix completed" - assert handoff["task_key"] == "TEST-1-ci-fix-1" - - def test_wait_for_ci_gate_does_not_exist(): """wait_for_ci_gate has been deleted from ci_evaluator module.""" import forge.workflow.nodes.ci_evaluator as mod diff --git a/tests/unit/workflow/nodes/test_code_review.py b/tests/unit/workflow/nodes/test_code_review.py index a13c9d626..7b39782ca 100644 --- a/tests/unit/workflow/nodes/test_code_review.py +++ b/tests/unit/workflow/nodes/test_code_review.py @@ -6,7 +6,6 @@ from forge.observability.review_poller import ReviewCycleData from forge.sandbox.runner import ContainerResult -from forge.integrations.github.client import PullRequestCreationResult from tests.fixtures.workflow_states import make_workflow_state FIX_COMMITS = ( @@ -34,10 +33,12 @@ async def test_commits_review_fixes_when_changes_exist(self): runner_mock = MagicMock() runner_mock.run = AsyncMock() - with patch("forge.workflow.nodes.code_review.ContainerRunner", return_value=runner_mock), \ - patch("forge.workflow.nodes.code_review.GitOperations", return_value=git_mock), \ - patch("forge.workflow.nodes.code_review.Workspace"), \ - patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"): + with ( + patch("forge.workflow.nodes.code_review.ContainerRunner", return_value=runner_mock), + patch("forge.workflow.nodes.code_review.GitOperations", return_value=git_mock), + patch("forge.workflow.nodes.code_review.Workspace"), + patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"), + ): committed, _ = await run_post_change_review( workspace_path="/tmp/ws", ticket_key="TEST-123", @@ -62,10 +63,12 @@ async def test_returns_false_when_no_changes(self): runner_mock = MagicMock() runner_mock.run = AsyncMock() - with patch("forge.workflow.nodes.code_review.ContainerRunner", return_value=runner_mock), \ - patch("forge.workflow.nodes.code_review.GitOperations", return_value=git_mock), \ - patch("forge.workflow.nodes.code_review.Workspace"), \ - patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"): + with ( + patch("forge.workflow.nodes.code_review.ContainerRunner", return_value=runner_mock), + patch("forge.workflow.nodes.code_review.GitOperations", return_value=git_mock), + patch("forge.workflow.nodes.code_review.Workspace"), + patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"), + ): committed, _ = await run_post_change_review( workspace_path="/tmp/ws", ticket_key="TEST-123", @@ -84,8 +87,10 @@ async def test_container_error_does_not_propagate(self): runner_mock = MagicMock() runner_mock.run = AsyncMock(side_effect=RuntimeError("container crashed")) - with patch("forge.workflow.nodes.code_review.ContainerRunner", return_value=runner_mock), \ - patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"): + with ( + patch("forge.workflow.nodes.code_review.ContainerRunner", return_value=runner_mock), + patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"), + ): committed, result = await run_post_change_review( workspace_path="/tmp/ws", ticket_key="TEST-123", @@ -144,7 +149,9 @@ async def test_returns_container_result_for_exhaustion_propagation(self): ) assert committed is False - assert container_result is not None, "ContainerResult must be returned for exhaustion propagation" + assert container_result is not None, ( + "ContainerResult must be returned for exhaustion propagation" + ) assert container_result.review_exhausted is True @@ -157,17 +164,42 @@ def _git_mock(commit_log: str = FIX_COMMITS) -> MagicMock: return git -def _github_jira_mocks(pr_body: str): - github = MagicMock() - github.get_pull_request = AsyncMock(return_value={"body": pr_body, "number": 42}) - github.update_pull_request = AsyncMock(return_value={"number": 42}) - github.close = AsyncMock() +def _adapter_jira_mocks(pr_body: str): + from forge.integrations.source_control.contracts import ( + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + ) + + adapter = AsyncMock() + adapter.get_change_request.return_value = ChangeRequest( + identity=ChangeRequestIdentity("c", "org/repo", 42), + url="u", + title="t", + body=pr_body, + state=ChangeRequestState.OPEN, + source_branch="f", + target_branch="main", + ) jira = MagicMock() jira.add_comment = AsyncMock() jira.close = AsyncMock() - return github, jira + return adapter, jira + + +def _repo_ref(): + from forge.integrations.source_control.contracts import Provider, RepositoryRef + + return RepositoryRef( + id="org/repo", + provider=Provider.GITHUB, + connection="c", + namespace="org/repo", + default_branch="main", + change_request_mode="fork", + ) class TestSyncPrDescription: @@ -184,23 +216,32 @@ async def test_updates_pr_when_description_is_inaccurate(self, state): original = "The jitter is +-10% uniform." updated = "The jitter is [0%, +20%] positive-only." - github, jira = _github_jira_mocks(original) + adapter, jira = _adapter_jira_mocks(original) agent_mock = MagicMock() agent_mock.run_task = AsyncMock(return_value=updated) agent_mock.close = AsyncMock() agent_mock._strip_preamble = MagicMock(side_effect=lambda x: x) - with patch("forge.workflow.nodes.code_review.GitHubClient", return_value=github), \ - patch("forge.workflow.nodes.code_review.JiraClient", return_value=jira), \ - patch("forge.workflow.nodes.code_review.ForgeAgent", return_value=agent_mock), \ - patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"): + with ( + patch( + "forge.workflow.nodes.code_review.get_adapter", return_value=(_repo_ref(), adapter) + ), + patch("forge.workflow.nodes.code_review.JiraClient", return_value=jira), + patch("forge.workflow.nodes.code_review.ForgeAgent", return_value=agent_mock), + patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"), + ): await sync_pr_description( - state, _git_mock(), - owner="org", repo="repo", pr_number=42, attempt=2, + state, + _git_mock(), + current_repo="org/repo", + pr_number=42, + attempt=2, ) - github.update_pull_request.assert_called_once_with("org", "repo", 42, body=updated) + adapter.update_change_request.assert_called_once() + _, kwargs = adapter.update_change_request.call_args + assert kwargs["body"] == updated jira.add_comment.assert_called_once() @pytest.mark.asyncio @@ -209,23 +250,30 @@ async def test_skips_when_body_unchanged(self, state): from forge.workflow.nodes.code_review import sync_pr_description body = "The jitter is +-10% uniform." - github, jira = _github_jira_mocks(body) + adapter, jira = _adapter_jira_mocks(body) agent_mock = MagicMock() agent_mock.run_task = AsyncMock(return_value=body) agent_mock.close = AsyncMock() agent_mock._strip_preamble = MagicMock(side_effect=lambda x: x) - with patch("forge.workflow.nodes.code_review.GitHubClient", return_value=github), \ - patch("forge.workflow.nodes.code_review.JiraClient", return_value=jira), \ - patch("forge.workflow.nodes.code_review.ForgeAgent", return_value=agent_mock), \ - patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"): + with ( + patch( + "forge.workflow.nodes.code_review.get_adapter", return_value=(_repo_ref(), adapter) + ), + patch("forge.workflow.nodes.code_review.JiraClient", return_value=jira), + patch("forge.workflow.nodes.code_review.ForgeAgent", return_value=agent_mock), + patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"), + ): await sync_pr_description( - state, _git_mock(), - owner="org", repo="repo", pr_number=42, attempt=2, + state, + _git_mock(), + current_repo="org/repo", + pr_number=42, + attempt=2, ) - github.update_pull_request.assert_not_called() + adapter.update_change_request.assert_not_called() jira.add_comment.assert_not_called() @pytest.mark.asyncio @@ -233,14 +281,21 @@ async def test_skips_when_no_commits(self, state): """Empty commit log skips the agent call entirely.""" from forge.workflow.nodes.code_review import sync_pr_description - github, jira = _github_jira_mocks("body") + adapter, jira = _adapter_jira_mocks("body") - with patch("forge.workflow.nodes.code_review.GitHubClient", return_value=github), \ - patch("forge.workflow.nodes.code_review.JiraClient", return_value=jira), \ - patch("forge.workflow.nodes.code_review.ForgeAgent") as MockAgent: + with ( + patch( + "forge.workflow.nodes.code_review.get_adapter", return_value=(_repo_ref(), adapter) + ), + patch("forge.workflow.nodes.code_review.JiraClient", return_value=jira), + patch("forge.workflow.nodes.code_review.ForgeAgent") as MockAgent, + ): await sync_pr_description( - state, _git_mock(""), - owner="org", repo="repo", pr_number=42, attempt=1, + state, + _git_mock(""), + current_repo="org/repo", + pr_number=42, + attempt=1, ) MockAgent.assert_not_called() @@ -250,54 +305,71 @@ async def test_skips_when_no_pr_number(self, state): """No PR number means nothing to update.""" from forge.workflow.nodes.code_review import sync_pr_description - with patch("forge.workflow.nodes.code_review.GitHubClient") as MockGH: + with patch("forge.workflow.nodes.code_review.get_adapter") as get_adapter_mock: await sync_pr_description( - state, MagicMock(), - owner="org", repo="repo", pr_number=None, attempt=1, + state, + MagicMock(), + current_repo="org/repo", + pr_number=None, + attempt=1, ) - MockGH.assert_not_called() + get_adapter_mock.assert_not_called() @pytest.mark.asyncio async def test_error_does_not_propagate(self, state): """Agent failure never blocks the caller.""" from forge.workflow.nodes.code_review import sync_pr_description - github, jira = _github_jira_mocks("body") + adapter, jira = _adapter_jira_mocks("body") agent_mock = MagicMock() agent_mock.run_task = AsyncMock(side_effect=RuntimeError("timeout")) agent_mock.close = AsyncMock() - with patch("forge.workflow.nodes.code_review.GitHubClient", return_value=github), \ - patch("forge.workflow.nodes.code_review.JiraClient", return_value=jira), \ - patch("forge.workflow.nodes.code_review.ForgeAgent", return_value=agent_mock), \ - patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"): + with ( + patch( + "forge.workflow.nodes.code_review.get_adapter", return_value=(_repo_ref(), adapter) + ), + patch("forge.workflow.nodes.code_review.JiraClient", return_value=jira), + patch("forge.workflow.nodes.code_review.ForgeAgent", return_value=agent_mock), + patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"), + ): await sync_pr_description( - state, _git_mock(), - owner="org", repo="repo", pr_number=42, attempt=1, + state, + _git_mock(), + current_repo="org/repo", + pr_number=42, + attempt=1, ) - github.update_pull_request.assert_not_called() + adapter.update_change_request.assert_not_called() @pytest.mark.asyncio async def test_audit_comment_labels_initial_create(self, state): """attempt=0 produces a human-readable 'PR creation' label in the comment.""" from forge.workflow.nodes.code_review import sync_pr_description - github, jira = _github_jira_mocks("old body") + adapter, jira = _adapter_jira_mocks("old body") agent_mock = MagicMock() agent_mock.run_task = AsyncMock(return_value="new body") agent_mock.close = AsyncMock() - with patch("forge.workflow.nodes.code_review.GitHubClient", return_value=github), \ - patch("forge.workflow.nodes.code_review.JiraClient", return_value=jira), \ - patch("forge.workflow.nodes.code_review.ForgeAgent", return_value=agent_mock), \ - patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"): + with ( + patch( + "forge.workflow.nodes.code_review.get_adapter", return_value=(_repo_ref(), adapter) + ), + patch("forge.workflow.nodes.code_review.JiraClient", return_value=jira), + patch("forge.workflow.nodes.code_review.ForgeAgent", return_value=agent_mock), + patch("forge.workflow.nodes.code_review.load_prompt", return_value="prompt"), + ): await sync_pr_description( - state, _git_mock(), - owner="org", repo="repo", pr_number=42, attempt=0, + state, + _git_mock(), + current_repo="org/repo", + pr_number=42, + attempt=0, ) comment_text = jira.add_comment.call_args[0][1] @@ -312,6 +384,14 @@ class TestSyncCalledFromCreatePR: @pytest.mark.asyncio async def test_sync_called_after_pr_creation(self): + from forge.integrations.source_control.contracts import ( + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + Provider, + RepositoryRef, + WriteTarget, + ) from forge.workflow.nodes.pr_creation import create_pull_request state = make_workflow_state( @@ -322,19 +402,54 @@ async def test_sync_called_after_pr_creation(self): context={"branch_name": "forge/test-123"}, ) - mock_github = MagicMock() - mock_github.get_or_create_fork = AsyncMock( - return_value={"owner": {"login": "fork-user"}, "name": "repo"} + repo_ref = RepositoryRef( + id="org/repo", + provider=Provider.GITHUB, + connection="c", + namespace="org/repo", + default_branch="main", + change_request_mode="fork", ) - mock_github.sync_fork_with_upstream = AsyncMock() - mock_github.add_fork_remote = MagicMock() - mock_github.create_pull_request = AsyncMock( - return_value=PullRequestCreationResult( - pr={"number": 42, "html_url": "https://github.com/org/repo/pull/42"}, + mock_adapter = AsyncMock() + mock_adapter.ensure_write_target = AsyncMock( + return_value=WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="", + base_branch="main", + fork_owner="fork-user", + fork_repo="repo", + ) + ) + mock_adapter.create_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="org/repo", native_id=42 + ), + url="https://github.com/org/repo/pull/42", + title="t", + body="b", + state=ChangeRequestState.OPEN, + source_branch="f", + target_branch="main", created=True, ) ) - mock_github.close = AsyncMock() + mock_adapter.get_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="org/repo", native_id=42 + ), + url="https://github.com/org/repo/pull/42", + title="t", + body="", + state=ChangeRequestState.OPEN, + source_branch="f", + target_branch="main", + ) + ) + mock_adapter.update_change_request = AsyncMock() + mock_adapter.create_comment = AsyncMock() mock_jira = MagicMock() mock_jira.get_issue = AsyncMock(return_value=MagicMock(summary="Test feature")) @@ -347,18 +462,27 @@ async def test_sync_called_after_pr_creation(self): mock_git.push_to_fork = MagicMock() mock_git.add_fork_remote = MagicMock() - with 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"), \ - patch("forge.workflow.nodes.pr_creation.check_merge_conflicts", - AsyncMock(return_value=(False, []))), \ - patch("forge.workflow.nodes.pr_creation._generate_pr_body_with_agent", - AsyncMock(return_value="## Summary\n\nTest PR.")), \ - patch("forge.workflow.nodes.pr_creation.set_pr_ticket_index", - new_callable=AsyncMock), \ - patch("forge.workflow.nodes.pr_creation.sync_pr_description", - new_callable=AsyncMock) as mock_sync: + with ( + patch( + "forge.workflow.nodes.pr_creation.get_adapter", + return_value=(repo_ref, mock_adapter), + ), + 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"), + patch( + "forge.workflow.nodes.pr_creation.check_merge_conflicts", + AsyncMock(return_value=(False, [])), + ), + patch( + "forge.workflow.nodes.pr_creation._generate_pr_body_with_agent", + AsyncMock(return_value="## Summary\n\nTest PR."), + ), + patch("forge.workflow.nodes.pr_creation.set_pr_ticket_index", new_callable=AsyncMock), + patch( + "forge.workflow.nodes.pr_creation.sync_pr_description", new_callable=AsyncMock + ) as mock_sync, + ): await create_pull_request(state) mock_sync.assert_called_once() diff --git a/tests/unit/workflow/nodes/test_human_review_gate.py b/tests/unit/workflow/nodes/test_human_review_gate.py index 6b527a332..5521afedf 100644 --- a/tests/unit/workflow/nodes/test_human_review_gate.py +++ b/tests/unit/workflow/nodes/test_human_review_gate.py @@ -39,6 +39,19 @@ def test_pending_ci_event_overrides_revision_requested(self): } assert route_human_review(state) == "ci_evaluator" + def test_pr_merged_takes_priority_over_pending_ci_event(self): + """A merge takes priority even when pending_ci_event is also set — + route_human_review must not send an already-merged PR through + ci_evaluator just because a stale/racing CI event is pending.""" + from forge.workflow.nodes.human_review import route_human_review + + state = { + **BASE_STATE, + "pending_ci_event": True, + "pr_merged": True, + } + assert route_human_review(state) == "complete_tasks" + def test_no_ci_event_paused_returns_end(self): """With no pending CI event, paused=True returns END.""" from forge.workflow.nodes.human_review import route_human_review diff --git a/tests/unit/workflow/nodes/test_local_reviewer.py b/tests/unit/workflow/nodes/test_local_reviewer.py index 1f5f8149c..d3c2fdf13 100644 --- a/tests/unit/workflow/nodes/test_local_reviewer.py +++ b/tests/unit/workflow/nodes/test_local_reviewer.py @@ -38,6 +38,8 @@ def base_bug_review_state(): "context": {"branch_name": "fix/BUG-42"}, "retry_count": 0, "last_error": None, + "fork_owner": "forge-bot", + "fork_repo": "backend", } @@ -55,6 +57,8 @@ def base_feature_review_state(): "context": {"branch_name": "feat/FEAT-10"}, "retry_count": 0, "last_error": None, + "fork_owner": "forge-bot", + "fork_repo": "backend", } diff --git a/tests/unit/workflow/nodes/test_model_policy_error_reporting.py b/tests/unit/workflow/nodes/test_model_policy_error_reporting.py index b0a1b9506..20d9d9857 100644 --- a/tests/unit/workflow/nodes/test_model_policy_error_reporting.py +++ b/tests/unit/workflow/nodes/test_model_policy_error_reporting.py @@ -4,6 +4,8 @@ import pytest +from forge.integrations.source_control.contracts import Provider, RepositoryRef +from forge.integrations.source_control.errors import NotFoundError from forge.workflow.nodes.error_handler import notify_error @@ -14,9 +16,15 @@ async def test_model_policy_error_is_reported_to_jira_and_active_github_pr() -> jira.get_issue = AsyncMock(return_value=issue) jira.add_model_policy_error_comment = AsyncMock() jira.close = AsyncMock() - github = MagicMock() - github.create_issue_comment = AsyncMock() - github.close = AsyncMock() + adapter = AsyncMock() + repo_ref = RepositoryRef( + id="forge-sdlc/forge", + provider=Provider.GITHUB, + connection="c", + namespace="forge-sdlc/forge", + default_branch="main", + change_request_mode="fork", + ) error = ( "Model policy configuration error: model 'missing' is not allowed on " "connection 'vertex'. Available connections and models: " @@ -30,7 +38,7 @@ async def test_model_policy_error_is_reported_to_jira_and_active_github_pr() -> with ( patch("forge.workflow.nodes.error_handler.JiraClient", return_value=jira), - patch("forge.integrations.github.client.GitHubClient", return_value=github), + patch("forge.workflow.nodes.error_handler.get_adapter", return_value=(repo_ref, adapter)), ): await notify_error(state, error, "implement_task") @@ -42,9 +50,41 @@ async def test_model_policy_error_is_reported_to_jira_and_active_github_pr() -> assert guidance["fix_command"] == ( "forge project-setup PROJ --model implement_task=CONNECTION:MODEL" ) - github.create_issue_comment.assert_awaited_once() - assert "vertex-ai: vertex=" in github.create_issue_comment.await_args.args[3] - github.close.assert_awaited_once() + adapter.create_comment.assert_awaited_once() + call_args = adapter.create_comment.call_args[0] + assert call_args[0] is repo_ref + assert "vertex-ai: vertex=" in call_args[2] + + +@pytest.mark.asyncio +async def test_github_mirror_failure_is_swallowed() -> None: + """A SourceControlError from get_adapter/create_comment doesn't fail notify_error.""" + jira = MagicMock() + jira.get_issue = AsyncMock(return_value=MagicMock(reporter=None, assignee=None)) + jira.add_model_policy_error_comment = AsyncMock() + jira.close = AsyncMock() + error = ( + "Model policy configuration error: model 'missing' is not allowed on " + "connection 'vertex'. Available connections and models: " + "vertex-ai: vertex=[gemini-3.5-flash, claude-sonnet-5]" + ) + state = { + "ticket_key": "PROJ-1", + "current_repo": "forge-sdlc/forge", + "current_pr_number": 251, + } + + with ( + patch("forge.workflow.nodes.error_handler.JiraClient", return_value=jira), + patch( + "forge.workflow.nodes.error_handler.get_adapter", + side_effect=NotFoundError("no adapter registered"), + ), + ): + await notify_error(state, error, "implement_task") + + jira.add_model_policy_error_comment.assert_awaited_once() + jira.close.assert_awaited_once() @pytest.mark.asyncio @@ -56,7 +96,7 @@ async def test_non_model_error_is_not_mirrored_to_github() -> None: with ( patch("forge.workflow.nodes.error_handler.JiraClient", return_value=jira), - patch("forge.integrations.github.client.GitHubClient") as github_type, + patch("forge.workflow.nodes.error_handler.get_adapter") as get_adapter_mock, ): await notify_error( { @@ -68,4 +108,4 @@ async def test_non_model_error_is_not_mirrored_to_github() -> None: "implement_task", ) - github_type.assert_not_called() + get_adapter_mock.assert_not_called() diff --git a/tests/unit/workflow/nodes/test_pr_creation_draft.py b/tests/unit/workflow/nodes/test_pr_creation_draft.py index 470b4be08..9e9befd65 100644 --- a/tests/unit/workflow/nodes/test_pr_creation_draft.py +++ b/tests/unit/workflow/nodes/test_pr_creation_draft.py @@ -1,38 +1,78 @@ """Unit tests for draft PR creation behavior and repository metadata configuration.""" from pathlib import Path -from unittest.mock import ANY, AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest -from forge.integrations.github.client import GitHubClient, PullRequestCreationResult +from forge.integrations.github.client import GitHubClient from forge.integrations.jira.client import JiraClient, MissingProjectConfig +from forge.integrations.source_control.contracts import ( + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + Provider, + RepositoryRef, + WriteTarget, +) from forge.workflow.feature.state import create_initial_feature_state from forge.workflow.nodes.pr_creation import create_pull_request -def create_mock_github_client(pr_number=123, pr_url="https://github.com/owner/repo/pull/123"): - """Create a mock GitHubClient with configurable PR data.""" - mock = MagicMock() - mock.close = AsyncMock() - mock.get_or_create_fork = AsyncMock( - return_value={ - "owner": {"login": "fork-owner"}, - "name": "repo", - } +def _repo_ref(identifier: str) -> RepositoryRef: + return RepositoryRef( + id=identifier, + provider=Provider.GITHUB, + connection="c", + namespace=identifier, + default_branch="main", + change_request_mode="fork", ) - mock.sync_fork_with_upstream = AsyncMock() - pr_data = { - "html_url": pr_url, - } - if pr_number is not None: - pr_data["number"] = pr_number - mock.create_pull_request = AsyncMock( - return_value=PullRequestCreationResult(pr=pr_data, created=True) +def create_mock_adapter(pr_number=123, pr_url="https://github.com/owner/repo/pull/123"): + """Create a mock SourceControlProvider adapter with configurable PR data.""" + adapter = AsyncMock() + adapter.ensure_write_target = AsyncMock( + return_value=WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="", + base_branch="main", + fork_owner="fork-owner", + fork_repo="repo", + ) ) - return mock + adapter.create_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="owner/repo", native_id=pr_number + ), + url=pr_url, + title="t", + body="b", + state=ChangeRequestState.OPEN, + source_branch="f", + target_branch="main", + created=True, + ) + ) + adapter.get_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="owner/repo", native_id=pr_number + ), + url=pr_url, + title="t", + body="", + state=ChangeRequestState.OPEN, + source_branch="f", + target_branch="main", + ) + ) + adapter.update_change_request = AsyncMock() + adapter.create_comment = AsyncMock() + return adapter def create_mock_jira_client(): @@ -244,8 +284,8 @@ class TestPRCreationNodeDraft: @pytest.mark.asyncio async def test_create_pr_node_invokes_github_with_draft_true(self): - """Workflow node should invoke GitHubClient with draft=True if repository configured as draft.""" - mock_github = create_mock_github_client(pr_number=101) + """Workflow node should invoke the adapter with draft=True if repository configured as draft.""" + mock_adapter = create_mock_adapter(pr_number=101) mock_jira = create_mock_jira_client() mock_jira.is_repo_draft = AsyncMock(return_value=True) mock_git = create_mock_git_operations() @@ -259,7 +299,10 @@ async def test_create_pr_node_invokes_github_with_draft_true(self): state["context"] = {"branch_name": "feat/test"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + patch( + "forge.workflow.nodes.pr_creation.get_adapter", + return_value=(_repo_ref("owner/repo"), mock_adapter), + ), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -281,21 +324,19 @@ async def test_create_pr_node_invokes_github_with_draft_true(self): # Verify is_repo_draft was resolved mock_jira.is_repo_draft.assert_called_once_with("FEAT", "owner/repo") - # Verify create_pull_request was called with draft=True - mock_github.create_pull_request.assert_called_once_with( - owner="owner", - repo="repo", - title="[FEAT-123] Test feature", - body=ANY, - head="fork-owner:feat/test", - base="main", - draft=True, - ) + # Verify create_change_request was called with draft=True and the fork head + mock_adapter.create_change_request.assert_called_once() + _, kwargs = mock_adapter.create_change_request.call_args + assert kwargs["title"] == "[FEAT-123] Test feature" + assert kwargs["target"].head_ref == "feat/test" + assert kwargs["target"].base_branch == "main" + assert kwargs["target"].fork_owner == "fork-owner" + assert kwargs["draft"] is True @pytest.mark.asyncio async def test_create_pr_node_invokes_github_with_draft_false(self): - """Workflow node should invoke GitHubClient with draft=False if repository is not configured as draft.""" - mock_github = create_mock_github_client(pr_number=102) + """Workflow node should invoke the adapter with draft=False if repository is not configured as draft.""" + mock_adapter = create_mock_adapter(pr_number=102) mock_jira = create_mock_jira_client() mock_jira.is_repo_draft = AsyncMock(return_value=False) mock_git = create_mock_git_operations() @@ -309,7 +350,10 @@ async def test_create_pr_node_invokes_github_with_draft_false(self): state["context"] = {"branch_name": "feat/test"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + patch( + "forge.workflow.nodes.pr_creation.get_adapter", + return_value=(_repo_ref("owner/repo"), mock_adapter), + ), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -331,13 +375,10 @@ async def test_create_pr_node_invokes_github_with_draft_false(self): # Verify is_repo_draft was resolved mock_jira.is_repo_draft.assert_called_once_with("FEAT", "owner/repo") - # Verify create_pull_request was called with draft=False - mock_github.create_pull_request.assert_called_once_with( - owner="owner", - repo="repo", - title="[FEAT-123] Test feature", - body=ANY, - head="fork-owner:feat/test", - base="main", - draft=False, - ) + # Verify create_change_request was called with draft=False + mock_adapter.create_change_request.assert_called_once() + _, kwargs = mock_adapter.create_change_request.call_args + assert kwargs["title"] == "[FEAT-123] Test feature" + assert kwargs["target"].head_ref == "feat/test" + assert kwargs["target"].base_branch == "main" + assert kwargs["draft"] is False diff --git a/tests/unit/workflow/nodes/test_pr_creation_informational_comment.py b/tests/unit/workflow/nodes/test_pr_creation_informational_comment.py index 92cc9aec4..5c472bdf8 100644 --- a/tests/unit/workflow/nodes/test_pr_creation_informational_comment.py +++ b/tests/unit/workflow/nodes/test_pr_creation_informational_comment.py @@ -5,37 +5,74 @@ import pytest +from forge.integrations.source_control.contracts import ( + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + Provider, + RepositoryRef, + WriteTarget, +) from forge.workflow.feature.state import create_initial_feature_state from forge.workflow.nodes.pr_creation import create_pull_request -from forge.integrations.github.client import PullRequestCreationResult -def create_mock_github_client( +def _repo_ref(identifier: str) -> RepositoryRef: + return RepositoryRef( + id=identifier, + provider=Provider.GITHUB, + connection="c", + namespace=identifier, + default_branch="main", + change_request_mode="fork", + ) + + +def create_mock_adapter( pr_number=123, pr_url="https://github.com/owner/repo/pull/123", is_new_pr=True ): - """Create a mock GitHubClient with configurable PR data.""" - mock = MagicMock() - mock.close = AsyncMock() - mock.get_or_create_fork = AsyncMock( - return_value={ - "owner": {"login": "fork-owner"}, - "name": "repo", - } + """Create a mock SourceControlProvider adapter with configurable PR data.""" + adapter = AsyncMock() + adapter.ensure_write_target = AsyncMock( + return_value=WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="", + base_branch="main", + fork_owner="fork-owner", + fork_repo="repo", + ) ) - mock.sync_fork_with_upstream = AsyncMock() - mock.create_issue_comment = AsyncMock() - - # PR creation response - can be configured for different scenarios - pr_data = { - "html_url": pr_url, - } - if pr_number is not None: - pr_data["number"] = pr_number - - mock.create_pull_request = AsyncMock( - return_value=PullRequestCreationResult(pr=pr_data, created=is_new_pr) + adapter.create_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="owner/repo", native_id=pr_number + ), + url=pr_url, + title="t", + body="b", + state=ChangeRequestState.OPEN, + source_branch="f", + target_branch="main", + created=is_new_pr, + ) ) - return mock + adapter.get_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="owner/repo", native_id=pr_number + ), + url=pr_url, + title="t", + body="", + state=ChangeRequestState.OPEN, + source_branch="f", + target_branch="main", + ) + ) + adapter.update_change_request = AsyncMock() + adapter.create_comment = AsyncMock() + return adapter def create_mock_jira_client(): @@ -91,13 +128,20 @@ def mock_external_pr_creation_side_effects(): yield +def _patch_adapter(adapter): + return patch( + "forge.workflow.nodes.pr_creation.get_adapter", + return_value=(_repo_ref("owner/repo"), adapter), + ) + + class TestPRInformationalComment: """Test cases for the informational PR command comment on PR creation.""" @pytest.mark.asyncio async def test_posts_comment_on_new_pr(self): """Should post an informational comment when a new PR is created.""" - mock_github = create_mock_github_client( + mock_adapter = create_mock_adapter( pr_number=456, pr_url="https://github.com/owner/repo/pull/456", is_new_pr=True ) mock_jira = create_mock_jira_client() @@ -112,7 +156,7 @@ async def test_posts_comment_on_new_pr(self): state["context"] = {"branch_name": "feat/test-branch"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -126,19 +170,17 @@ async def test_posts_comment_on_new_pr(self): await create_pull_request(state) # Verify comment was posted - mock_github.create_issue_comment.assert_called_once() - call_args = mock_github.create_issue_comment.call_args - assert call_args[1]["owner"] == "owner" - assert call_args[1]["repo"] == "repo" - assert call_args[1]["issue_number"] == 456 - assert "/forge rebase" in call_args[1]["body"] - assert "/forge skip-gate" in call_args[1]["body"] - assert "/forge unskip-gate" in call_args[1]["body"] + mock_adapter.create_comment.assert_awaited_once() + call_args = mock_adapter.create_comment.call_args[0] + assert call_args[1].native_id == 456 + assert "/forge rebase" in call_args[2] + assert "/forge skip-gate" in call_args[2] + assert "/forge unskip-gate" in call_args[2] @pytest.mark.asyncio async def test_does_not_post_comment_on_existing_pr(self): """Should NOT post an informational comment when an existing PR is returned.""" - mock_github = create_mock_github_client( + mock_adapter = create_mock_adapter( pr_number=456, pr_url="https://github.com/owner/repo/pull/456", is_new_pr=False ) mock_jira = create_mock_jira_client() @@ -153,7 +195,7 @@ async def test_does_not_post_comment_on_existing_pr(self): state["context"] = {"branch_name": "feat/test-branch"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -167,16 +209,16 @@ async def test_does_not_post_comment_on_existing_pr(self): await create_pull_request(state) # Verify comment was NOT posted - mock_github.create_issue_comment.assert_not_called() + mock_adapter.create_comment.assert_not_called() @pytest.mark.asyncio async def test_comment_failure_ignored(self): """Should successfully finish PR creation even if posting comment fails.""" - mock_github = create_mock_github_client( + mock_adapter = create_mock_adapter( pr_number=456, pr_url="https://github.com/owner/repo/pull/456", is_new_pr=True ) # Make the comment creation raise an exception - mock_github.create_issue_comment.side_effect = Exception("GitHub API Error") + mock_adapter.create_comment.side_effect = Exception("GitHub API Error") mock_jira = create_mock_jira_client() mock_git = create_mock_git_operations() @@ -189,7 +231,7 @@ async def test_comment_failure_ignored(self): state["context"] = {"branch_name": "feat/test-branch"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( 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 68060510b..b80eec9f4 100644 --- a/tests/unit/workflow/nodes/test_pr_creation_pr_number.py +++ b/tests/unit/workflow/nodes/test_pr_creation_pr_number.py @@ -5,34 +5,72 @@ import pytest +from forge.integrations.source_control.contracts import ( + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + Provider, + RepositoryRef, + WriteTarget, +) from forge.workflow.feature.state import create_initial_feature_state from forge.workflow.nodes.pr_creation import create_pull_request -from forge.integrations.github.client import PullRequestCreationResult -def create_mock_github_client(pr_number=123, pr_url="https://github.com/owner/repo/pull/123"): - """Create a mock GitHubClient with configurable PR data.""" - mock = MagicMock() - mock.close = AsyncMock() - mock.get_or_create_fork = AsyncMock( - return_value={ - "owner": {"login": "fork-owner"}, - "name": "repo", - } +def _repo_ref(identifier: str) -> RepositoryRef: + return RepositoryRef( + id=identifier, + provider=Provider.GITHUB, + connection="c", + namespace=identifier, + default_branch="main", + change_request_mode="fork", ) - mock.sync_fork_with_upstream = AsyncMock() - # PR creation response - can be configured for different scenarios - pr_data = { - "html_url": pr_url, - } - if pr_number is not None: - pr_data["number"] = pr_number - mock.create_pull_request = AsyncMock( - return_value=PullRequestCreationResult(pr=pr_data, created=True) +def create_mock_adapter(pr_number=123, pr_url="https://github.com/owner/repo/pull/123"): + """Create a mock SourceControlProvider adapter with configurable PR data.""" + adapter = AsyncMock() + adapter.ensure_write_target = AsyncMock( + return_value=WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="", + base_branch="main", + fork_owner="fork-owner", + fork_repo="repo", + ) ) - return mock + adapter.create_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="owner/repo", native_id=pr_number + ), + url=pr_url, + title="t", + body="b", + state=ChangeRequestState.OPEN, + source_branch="f", + target_branch="main", + created=True, + ) + ) + adapter.get_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="owner/repo", native_id=pr_number + ), + url=pr_url, + title="t", + body="", + state=ChangeRequestState.OPEN, + source_branch="f", + target_branch="main", + ) + ) + adapter.update_change_request = AsyncMock() + adapter.create_comment = AsyncMock() + return adapter def create_mock_jira_client(): @@ -88,13 +126,20 @@ def mock_external_pr_creation_side_effects(): yield +def _patch_adapter(adapter, identifier="owner/repo"): + return patch( + "forge.workflow.nodes.pr_creation.get_adapter", + return_value=(_repo_ref(identifier), adapter), + ) + + class TestPRNumberExtractionSuccess: """Test cases for successful PR number extraction from GitHub API response.""" @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( + mock_adapter = create_mock_adapter( pr_number=456, pr_url="https://github.com/owner/repo/pull/456" ) mock_jira = create_mock_jira_client() @@ -109,7 +154,7 @@ async def test_pr_number_extracted_from_github_response(self): state["context"] = {"branch_name": "feat/test-branch"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -128,7 +173,7 @@ async def test_pr_number_extracted_from_github_response(self): @pytest.mark.asyncio async def test_pr_number_used_in_jira_remote_link(self): """Should use PR number in Jira remote link label when available.""" - mock_github = create_mock_github_client(pr_number=789) + mock_adapter = create_mock_adapter(pr_number=789) mock_jira = create_mock_jira_client() mock_git = create_mock_git_operations() @@ -141,7 +186,7 @@ async def test_pr_number_used_in_jira_remote_link(self): state["context"] = {"branch_name": "feat/test"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -162,7 +207,7 @@ async def test_pr_number_used_in_jira_remote_link(self): @pytest.mark.asyncio async def test_pr_number_used_in_info_logging(self, caplog): """Should include PR number in info log message when available.""" - mock_github = create_mock_github_client(pr_number=999) + mock_adapter = create_mock_adapter(pr_number=999) mock_jira = create_mock_jira_client() mock_git = create_mock_git_operations() @@ -175,7 +220,7 @@ async def test_pr_number_used_in_info_logging(self, caplog): state["context"] = {"branch_name": "feat/test"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -196,6 +241,54 @@ async def test_pr_number_used_in_info_logging(self, caplog): ) +class TestPRCreationDirectMode: + """Direct-mode repos (no fork identity) must push to origin, not 'fork'.""" + + @pytest.mark.asyncio + async def test_direct_mode_pushes_to_origin_not_fork(self): + mock_adapter = create_mock_adapter(pr_number=1) + mock_adapter.ensure_write_target = AsyncMock( + return_value=WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="", + base_branch="main", + fork_owner=None, + fork_repo=None, + ) + ) + mock_jira = create_mock_jira_client() + mock_git = create_mock_git_operations() + mock_git.push = MagicMock() + + state = create_initial_feature_state( + ticket_key="FEAT-DIRECT", + current_repo="owner/repo", + ) + state["workspace_path"] = "/tmp/test-workspace" + state["implemented_tasks"] = ["TASK-1"] + state["context"] = {"branch_name": "feat/test"} + + with ( + _patch_adapter(mock_adapter), + 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.sync_pr_description", new_callable=AsyncMock), + ): + result = await create_pull_request(state) + + mock_git.add_fork_remote.assert_not_called() + mock_git.push_to_fork.assert_not_called() + mock_git.push.assert_called_once_with(force=False) + assert result["last_error"] is None + + class TestPRNumberExtractionMissing: """Test cases for handling missing PR number in GitHub API response.""" @@ -203,7 +296,7 @@ class TestPRNumberExtractionMissing: async def test_pr_number_none_when_unavailable(self): """Should set current_pr_number to None when PR number unavailable in API response.""" # GitHub API returns response without 'number' field - mock_github = create_mock_github_client(pr_number=None) + mock_adapter = create_mock_adapter(pr_number=None) mock_jira = create_mock_jira_client() mock_git = create_mock_git_operations() @@ -216,7 +309,7 @@ async def test_pr_number_none_when_unavailable(self): state["context"] = {"branch_name": "feat/test"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -229,15 +322,17 @@ async def test_pr_number_none_when_unavailable(self): ): result = await create_pull_request(state) - # Verify PR number is None in state + # Verify PR number is None in state. With no number known, the record is + # keyed by URL (f"{repo}:{url}") — see workflow.pr_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"] + url_key = f"owner/repo:{result['current_pr_url']}" + assert result["pull_requests"][url_key]["number"] is None + assert result["pull_requests"][url_key]["url"] == result["current_pr_url"] @pytest.mark.asyncio async def test_workflow_continues_when_pr_number_unavailable(self): """Should continue workflow successfully even when PR number unavailable.""" - mock_github = create_mock_github_client(pr_number=None) + mock_adapter = create_mock_adapter(pr_number=None) mock_jira = create_mock_jira_client() mock_git = create_mock_git_operations() @@ -250,7 +345,7 @@ async def test_workflow_continues_when_pr_number_unavailable(self): state["context"] = {"branch_name": "feat/test"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -275,7 +370,7 @@ async def test_workflow_continues_when_pr_number_unavailable(self): async def test_warning_logged_when_pr_number_unavailable(self, caplog): """Should log warning with diagnostic info when PR number unavailable.""" pr_url = "https://github.com/owner/repo/pull/123" - mock_github = create_mock_github_client(pr_number=None, pr_url=pr_url) + mock_adapter = create_mock_adapter(pr_number=None, pr_url=pr_url) mock_jira = create_mock_jira_client() mock_git = create_mock_git_operations() @@ -288,7 +383,7 @@ async def test_warning_logged_when_pr_number_unavailable(self, caplog): state["context"] = {"branch_name": "feat/test"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -313,7 +408,7 @@ async def test_warning_logged_when_pr_number_unavailable(self, caplog): @pytest.mark.asyncio async def test_generic_label_used_when_pr_number_unavailable(self): """Should use generic 'Pull Request' label in Jira remote link when PR number unavailable.""" - mock_github = create_mock_github_client(pr_number=None) + mock_adapter = create_mock_adapter(pr_number=None) mock_jira = create_mock_jira_client() mock_git = create_mock_git_operations() @@ -326,7 +421,7 @@ async def test_generic_label_used_when_pr_number_unavailable(self): state["context"] = {"branch_name": "feat/test"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -348,7 +443,7 @@ async def test_generic_label_used_when_pr_number_unavailable(self): async def test_info_log_indicates_number_unavailable(self, caplog): """Should log info message indicating PR number unavailable.""" pr_url = "https://github.com/owner/repo/pull/456" - mock_github = create_mock_github_client(pr_number=None, pr_url=pr_url) + mock_adapter = create_mock_adapter(pr_number=None, pr_url=pr_url) mock_jira = create_mock_jira_client() mock_git = create_mock_git_operations() @@ -361,7 +456,7 @@ async def test_info_log_indicates_number_unavailable(self, caplog): state["context"] = {"branch_name": "feat/test"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -389,7 +484,7 @@ class TestPRNumberExtractionEdgeCases: async def test_pr_number_zero_handled_correctly(self): """Should handle PR number 0 (edge case) correctly without treating it as None.""" # PR number 0 is technically valid (though rare) and should not be treated as missing - mock_github = create_mock_github_client(pr_number=0) + mock_adapter = create_mock_adapter(pr_number=0) mock_jira = create_mock_jira_client() mock_git = create_mock_git_operations() @@ -402,7 +497,7 @@ async def test_pr_number_zero_handled_correctly(self): state["context"] = {"branch_name": "feat/test"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -427,7 +522,7 @@ async def test_pr_number_zero_handled_correctly(self): async def test_pr_number_extracted_when_pr_url_missing(self): """Should extract PR number even when PR URL is missing from response.""" # Edge case: API returns number but not html_url - mock_github = create_mock_github_client(pr_number=111, pr_url="") + mock_adapter = create_mock_adapter(pr_number=111, pr_url="") mock_jira = create_mock_jira_client() mock_git = create_mock_git_operations() @@ -440,7 +535,7 @@ async def test_pr_number_extracted_when_pr_url_missing(self): state["context"] = {"branch_name": "feat/test"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github), + _patch_adapter(mock_adapter), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -460,7 +555,7 @@ async def test_pr_number_extracted_when_pr_url_missing(self): async def test_multiple_prs_each_have_own_pr_number(self): """Should handle multiple PR creations with different PR numbers independently.""" # This tests that pr_number is properly isolated per PR creation - mock_github_1 = create_mock_github_client(pr_number=100) + mock_adapter_1 = create_mock_adapter(pr_number=100) mock_jira = create_mock_jira_client() mock_git = create_mock_git_operations() @@ -473,7 +568,7 @@ async def test_multiple_prs_each_have_own_pr_number(self): state["context"] = {"branch_name": "feat/test"} with ( - patch("forge.workflow.nodes.pr_creation.GitHubClient", return_value=mock_github_1), + _patch_adapter(mock_adapter_1), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -488,16 +583,16 @@ async def test_multiple_prs_each_have_own_pr_number(self): # Verify first PR has correct number assert result_1["current_pr_number"] == 100 - assert result_1["pull_requests"]["owner/repo"]["number"] == 100 + assert result_1["pull_requests"]["owner/repo:100"]["number"] == 100 # Simulate second PR creation with different number - mock_github_2 = create_mock_github_client( + mock_adapter_2 = create_mock_adapter( 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_adapter(mock_adapter_2, identifier="owner/other"), patch("forge.workflow.nodes.pr_creation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.pr_creation.GitOperations", return_value=mock_git), patch( @@ -510,7 +605,7 @@ async def test_multiple_prs_each_have_own_pr_number(self): ): result_2 = await create_pull_request(result_1) - # Verify second PR has correct number + # Verify second PR has correct number; each PR occupies its own slot. 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 + assert result_2["pull_requests"]["owner/repo:100"]["number"] == 100 + assert result_2["pull_requests"]["owner/other:200"]["number"] == 200 diff --git a/tests/unit/workflow/nodes/test_prd_pr.py b/tests/unit/workflow/nodes/test_prd_pr.py index 81d0580c8..cad352124 100644 --- a/tests/unit/workflow/nodes/test_prd_pr.py +++ b/tests/unit/workflow/nodes/test_prd_pr.py @@ -4,9 +4,39 @@ import pytest +from forge.integrations.source_control.contracts import ( + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + Provider, + RepositoryRef, + WriteTarget, +) +from forge.integrations.source_control.errors import ConflictError from forge.models.workflow import TicketType from forge.workflow.feature.state import create_initial_feature_state -from forge.integrations.github.client import PullRequestCreationResult + + +def _repo_ref(identifier: str = "org/proposals") -> RepositoryRef: + return RepositoryRef( + id=identifier, + provider=Provider.GITHUB, + connection="c", + namespace=identifier, + default_branch="main", + change_request_mode="fork", + ) + + +def _fork_ref(fork_owner: str, fork_repo: str) -> RepositoryRef: + return RepositoryRef( + id=f"{fork_owner}/{fork_repo}", + provider=Provider.GITHUB, + connection="c", + namespace=f"{fork_owner}/{fork_repo}", + default_branch="main", + change_request_mode="direct", + ) class TestCreatePrdProposalPr: @@ -14,29 +44,33 @@ class TestCreatePrdProposalPr: async def test_creates_branch_and_pr(self): from forge.workflow.nodes.prd_generation import _create_prd_proposal_pr - mock_gh = MagicMock() - mock_gh.get_or_create_fork = AsyncMock( - return_value={ - "owner": {"login": "forge-bot"}, - "name": "proposals", - "default_branch": "trunk", - } + mock_adapter = AsyncMock() + mock_adapter.resolve_default_branch = AsyncMock(return_value="trunk") + mock_adapter.ensure_write_target = AsyncMock( + return_value=WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="", + base_branch="main", + fork_owner="forge-bot", + fork_repo="proposals", + ) ) - mock_gh.get_repository = AsyncMock(return_value={"default_branch": "trunk"}) - mock_gh.sync_fork_with_upstream = AsyncMock(return_value=True) - mock_gh.get_file_contents = AsyncMock(return_value={"sha": "existing-sha"}) - mock_gh.create_branch = AsyncMock(return_value={"ref": "refs/heads/forge/prd/test-123"}) - mock_gh.create_or_update_file = AsyncMock(return_value={"content": {"sha": "filesha"}}) - mock_gh.create_pull_request = AsyncMock( - return_value=PullRequestCreationResult( - pr={ - "number": 7, - "html_url": "https://github.com/org/proposals/pull/7", - }, - created=True, + mock_adapter.create_branch = AsyncMock() + mock_adapter.put_file = AsyncMock() + mock_adapter.create_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="org/proposals", native_id=7 + ), + url="https://github.com/org/proposals/pull/7", + title="t", + body="b", + state=ChangeRequestState.OPEN, + source_branch="forge/prd/test-123", + target_branch="trunk", ) ) - mock_gh.close = AsyncMock() mock_jira = MagicMock() mock_jira.add_comment = AsyncMock() @@ -44,7 +78,10 @@ async def test_creates_branch_and_pr(self): mock_jira.close = AsyncMock() with ( - patch("forge.workflow.nodes.proposal_pr.GitHubClient", return_value=mock_gh), + patch( + "forge.workflow.nodes.proposal_pr.get_adapter", + return_value=(_repo_ref(), mock_adapter), + ), patch("forge.workflow.nodes.proposal_pr.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.proposal_pr.set_pr_ticket_index", @@ -61,54 +98,126 @@ async def test_creates_branch_and_pr(self): assert result["prd_pr_number"] == 7 assert result["prd_pr_url"] == "https://github.com/org/proposals/pull/7" assert result["prd_pr_repo"] == "org/proposals" - assert result["prd_pr_branch"] == "forge/prd/test-123" - assert result["prd_pr_file_path"] == "TEST-123/prd.md" - mock_gh.create_branch.assert_called_once_with( - "forge-bot", "proposals", "forge/prd/test-123", base="trunk" - ) - mock_gh.sync_fork_with_upstream.assert_awaited_once_with( - "forge-bot", "proposals", branch="trunk" + mock_adapter.create_branch.assert_called_once_with( + _fork_ref("forge-bot", "proposals"), "forge/prd/test-123", "trunk" ) - assert result["prd_pr_fork_owner"] == "forge-bot" - assert result["prd_pr_fork_repo"] == "proposals" - assert mock_gh.create_or_update_file.call_args.kwargs["sha"] == "existing-sha" - mock_gh.create_pull_request.assert_called_once() - pr_call_kwargs = mock_gh.create_pull_request.call_args[1] - assert pr_call_kwargs["owner"] == "org" - assert pr_call_kwargs["repo"] == "proposals" - assert pr_call_kwargs["head"] == "forge-bot:forge/prd/test-123" - assert pr_call_kwargs["base"] == "trunk" - assert "# My PRD" not in pr_call_kwargs["body"] - assert "TEST-123/prd.md" in pr_call_kwargs["body"] + mock_adapter.put_file.assert_called_once() + put_args = mock_adapter.put_file.call_args[0] + assert put_args[0] == _fork_ref("forge-bot", "proposals") + assert put_args[1] == "TEST-123/prd.md" + assert put_args[2] == "# My PRD" + assert put_args[4] == "forge/prd/test-123" + + mock_adapter.create_change_request.assert_called_once() + cr_args, cr_kwargs = mock_adapter.create_change_request.call_args + assert cr_args[0] == _repo_ref() + assert cr_args[1].head_ref == "forge/prd/test-123" + assert cr_args[1].base_branch == "trunk" + assert cr_args[1].fork_owner == "forge-bot" + assert "# My PRD" not in cr_kwargs["body"] + assert "TEST-123/prd.md" in cr_kwargs["body"] + mock_jira.add_comment.assert_called_once() mock_jira.set_workflow_label.assert_called_once() mock_index.assert_called_once() + @pytest.mark.asyncio + async def test_stores_canonical_namespace_not_repos_yaml_alias(self): + """prd_pr_repo must be the canonical namespace, not the raw (possibly + repos.yaml-alias) proposals_repo -- webhook matching in worker.py's + _is_prd_pr_event compares this against event.repo_ref.namespace, + which is always canonical.""" + from forge.workflow.nodes.prd_generation import _create_prd_proposal_pr + + mock_adapter = AsyncMock() + mock_adapter.resolve_default_branch = AsyncMock(return_value="main") + mock_adapter.ensure_write_target = AsyncMock( + return_value=WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="", + base_branch="main", + fork_owner="forge-bot", + fork_repo="proposals", + ) + ) + mock_adapter.create_branch = AsyncMock() + mock_adapter.put_file = AsyncMock() + mock_adapter.create_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="org/proposals", native_id=7 + ), + url="https://github.com/org/proposals/pull/7", + title="t", + body="b", + state=ChangeRequestState.OPEN, + source_branch="forge/prd/test-123", + target_branch="main", + ) + ) + + mock_jira = MagicMock() + mock_jira.add_comment = AsyncMock() + mock_jira.set_workflow_label = AsyncMock() + mock_jira.close = AsyncMock() + + with ( + patch( + "forge.workflow.nodes.proposal_pr.get_adapter", + return_value=(_repo_ref("org/proposals"), mock_adapter), + ), + patch("forge.workflow.nodes.proposal_pr.JiraClient", return_value=mock_jira), + patch( + "forge.workflow.nodes.proposal_pr.set_pr_ticket_index", + new_callable=AsyncMock, + ), + ): + result = await _create_prd_proposal_pr( + ticket_key="TEST-123", + prd_content="# My PRD", + summary="My Feature", + proposals_repo="proposals-alias", + ) + + assert result["prd_pr_repo"] == "org/proposals" + assert result["prd_pr_branch"] == "forge/prd/test-123" + assert result["prd_pr_file_path"] == "TEST-123/prd.md" + assert result["prd_pr_fork_owner"] == "forge-bot" + assert result["prd_pr_fork_repo"] == "proposals" + @pytest.mark.asyncio async def test_creates_pr_with_custom_path(self): from forge.workflow.nodes.prd_generation import _create_prd_proposal_pr - mock_gh = MagicMock() - mock_gh.get_or_create_fork = AsyncMock( - return_value={"owner": {"login": "forge-bot"}, "name": "proposals"} + mock_adapter = AsyncMock() + mock_adapter.resolve_default_branch = AsyncMock(return_value="main") + mock_adapter.ensure_write_target = AsyncMock( + return_value=WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="", + base_branch="main", + fork_owner="forge-bot", + fork_repo="proposals", + ) ) - # Omitting a custom default branch exercises the upstream "main" fallback. - mock_gh.get_repository = AsyncMock(return_value={}) - mock_gh.sync_fork_with_upstream = AsyncMock(return_value=True) - mock_gh.get_file_contents = AsyncMock(return_value=None) - mock_gh.create_branch = AsyncMock(return_value={"ref": "refs/heads/forge/prd/test-456"}) - mock_gh.create_or_update_file = AsyncMock(return_value={"content": {"sha": "filesha"}}) - mock_gh.create_pull_request = AsyncMock( - return_value=PullRequestCreationResult( - pr={ - "number": 10, - "html_url": "https://github.com/org/proposals/pull/10", - }, - created=True, + mock_adapter.create_branch = AsyncMock() + mock_adapter.put_file = AsyncMock() + mock_adapter.create_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="org/proposals", native_id=10 + ), + url="https://github.com/org/proposals/pull/10", + title="t", + body="b", + state=ChangeRequestState.OPEN, + source_branch="forge/prd/test-456", + target_branch="main", ) ) - mock_gh.close = AsyncMock() mock_jira = MagicMock() mock_jira.add_comment = AsyncMock() @@ -116,7 +225,10 @@ async def test_creates_pr_with_custom_path(self): mock_jira.close = AsyncMock() with ( - patch("forge.workflow.nodes.proposal_pr.GitHubClient", return_value=mock_gh), + patch( + "forge.workflow.nodes.proposal_pr.get_adapter", + return_value=(_repo_ref(), mock_adapter), + ), patch("forge.workflow.nodes.proposal_pr.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.proposal_pr.set_pr_ticket_index", @@ -132,28 +244,28 @@ async def test_creates_pr_with_custom_path(self): ) assert result["prd_pr_file_path"] == "enhancements/TEST-456/prd.md" - pr_call_kwargs = mock_gh.create_pull_request.call_args[1] - assert "enhancements/TEST-456/prd.md" in pr_call_kwargs["body"] - assert pr_call_kwargs["base"] == "main" - assert mock_gh.create_or_update_file.call_args.kwargs["sha"] is None + cr_args, cr_kwargs = mock_adapter.create_change_request.call_args + assert "enhancements/TEST-456/prd.md" in cr_kwargs["body"] + assert cr_args[1].base_branch == "main" @pytest.mark.asyncio async def test_stops_when_fork_cannot_be_synchronized(self): from forge.workflow.nodes.prd_generation import _create_prd_proposal_pr - mock_gh = MagicMock() - mock_gh.get_or_create_fork = AsyncMock( - return_value={"owner": {"login": "forge-bot"}, "name": "proposals"} + mock_adapter = AsyncMock() + mock_adapter.resolve_default_branch = AsyncMock(return_value="main") + mock_adapter.ensure_write_target = AsyncMock( + side_effect=ConflictError("Could not synchronize proposal fork") ) - mock_gh.get_repository = AsyncMock(return_value={"default_branch": "main"}) - mock_gh.sync_fork_with_upstream = AsyncMock(return_value=False) - mock_gh.close = AsyncMock() mock_jira = MagicMock(close=AsyncMock()) with ( - patch("forge.workflow.nodes.proposal_pr.GitHubClient", return_value=mock_gh), + patch( + "forge.workflow.nodes.proposal_pr.get_adapter", + return_value=(_repo_ref(), mock_adapter), + ), patch("forge.workflow.nodes.proposal_pr.JiraClient", return_value=mock_jira), - pytest.raises(RuntimeError, match="Could not synchronize proposal fork"), + pytest.raises(ConflictError, match="Could not synchronize proposal fork"), ): await _create_prd_proposal_pr( ticket_key="TEST-FAIL", @@ -162,9 +274,9 @@ async def test_stops_when_fork_cannot_be_synchronized(self): proposals_repo="org/proposals", ) - mock_gh.create_branch.assert_not_called() - mock_gh.create_or_update_file.assert_not_called() - mock_gh.create_pull_request.assert_not_called() + mock_adapter.create_branch.assert_not_called() + mock_adapter.put_file.assert_not_called() + mock_adapter.create_change_request.assert_not_called() class TestResolveProposalsPath: @@ -188,13 +300,10 @@ class TestUpdatePrdProposalPr: async def test_updates_file_on_branch(self): from forge.workflow.nodes.prd_generation import _update_prd_proposal_pr - mock_gh = MagicMock() - mock_gh.get_file_contents = AsyncMock( - return_value={"sha": "oldsha", "path": "TEST-123/prd.md"} - ) - mock_gh.create_or_update_file = AsyncMock(return_value={"content": {"sha": "newsha"}}) - mock_gh.create_issue_comment = AsyncMock() - mock_gh.close = AsyncMock() + mock_adapter = AsyncMock() + mock_adapter.get_file = AsyncMock(return_value="# Old PRD") + mock_adapter.put_file = AsyncMock() + mock_adapter.create_comment = AsyncMock() state = create_initial_feature_state( ticket_key="TEST-123", @@ -208,36 +317,42 @@ async def test_updates_file_on_branch(self): prd_pr_file_path="TEST-123/prd.md", ) - with patch("forge.workflow.nodes.proposal_pr.GitHubClient", return_value=mock_gh): + with patch( + "forge.workflow.nodes.proposal_pr.get_adapter", + return_value=(_repo_ref(), mock_adapter), + ): await _update_prd_proposal_pr( ticket_key="TEST-123", prd_content="# Revised PRD", state=state, ) - mock_gh.get_file_contents.assert_called_once_with( - "forge-bot", "proposals", "TEST-123/prd.md", "forge/prd/test-123" + mock_adapter.get_file.assert_called_once_with( + _fork_ref("forge-bot", "proposals"), "TEST-123/prd.md", "forge/prd/test-123" ) - mock_gh.create_or_update_file.assert_called_once() - call_kwargs = mock_gh.create_or_update_file.call_args[1] - assert call_kwargs["sha"] == "oldsha" - assert call_kwargs["path"] == "TEST-123/prd.md" - mock_gh.create_issue_comment.assert_called_once_with( - "org", - "proposals", - 7, - "PRD has been revised based on feedback. Please review the updated version.", + mock_adapter.put_file.assert_called_once() + put_args = mock_adapter.put_file.call_args[0] + assert put_args[0] == _fork_ref("forge-bot", "proposals") + assert put_args[1] == "TEST-123/prd.md" + assert put_args[2] == "# Revised PRD" + mock_adapter.create_comment.assert_called_once() + comment_args = mock_adapter.create_comment.call_args[0] + assert comment_args[0] == _repo_ref() + assert comment_args[1].native_id == 7 + assert ( + comment_args[2] + == "PRD has been revised based on feedback. Please review the updated version." ) @pytest.mark.asyncio async def test_updates_legacy_upstream_branch_without_fork_state(self): from forge.workflow.nodes.prd_generation import _update_prd_proposal_pr - mock_gh = MagicMock() - mock_gh.get_file_contents = AsyncMock(return_value={"sha": "oldsha"}) - mock_gh.create_or_update_file = AsyncMock() - mock_gh.create_issue_comment = AsyncMock() - mock_gh.close = AsyncMock() + mock_adapter = AsyncMock() + mock_adapter.get_file = AsyncMock(return_value="# Old") + mock_adapter.put_file = AsyncMock() + mock_adapter.create_comment = AsyncMock() + state = create_initial_feature_state( ticket_key="TEST-LEGACY", ticket_type=TicketType.FEATURE, @@ -247,30 +362,28 @@ async def test_updates_legacy_upstream_branch_without_fork_state(self): prd_pr_file_path="TEST-LEGACY/prd.md", ) - with patch("forge.workflow.nodes.proposal_pr.GitHubClient", return_value=mock_gh): + with patch( + "forge.workflow.nodes.proposal_pr.get_adapter", + return_value=(_repo_ref(), mock_adapter), + ): await _update_prd_proposal_pr("TEST-LEGACY", "# Revised", state) - mock_gh.get_file_contents.assert_awaited_once_with( - "org", "proposals", "TEST-LEGACY/prd.md", "forge/prd/test-legacy" + mock_adapter.get_file.assert_awaited_once_with( + _repo_ref(), "TEST-LEGACY/prd.md", "forge/prd/test-legacy" ) - assert mock_gh.create_or_update_file.call_args.kwargs["owner"] == "org" - assert mock_gh.create_or_update_file.call_args.kwargs["repo"] == "proposals" + put_args = mock_adapter.put_file.call_args[0] + assert put_args[0] == _repo_ref() @pytest.mark.asyncio async def test_skips_commit_when_content_unchanged(self): from forge.workflow.nodes.prd_generation import _update_prd_proposal_pr - from forge.workflow.nodes.proposal_pr import _git_blob_sha unchanged_content = "# Same PRD" - matching_sha = _git_blob_sha(unchanged_content) - mock_gh = MagicMock() - mock_gh.get_file_contents = AsyncMock( - return_value={"sha": matching_sha, "path": "TEST-123/prd.md"} - ) - mock_gh.create_or_update_file = AsyncMock() - mock_gh.create_issue_comment = AsyncMock() - mock_gh.close = AsyncMock() + mock_adapter = AsyncMock() + mock_adapter.get_file = AsyncMock(return_value=unchanged_content) + mock_adapter.put_file = AsyncMock() + mock_adapter.create_comment = AsyncMock() state = create_initial_feature_state( ticket_key="TEST-123", @@ -284,14 +397,17 @@ async def test_skips_commit_when_content_unchanged(self): prd_pr_file_path="TEST-123/prd.md", ) - with patch("forge.workflow.nodes.proposal_pr.GitHubClient", return_value=mock_gh): + with patch( + "forge.workflow.nodes.proposal_pr.get_adapter", + return_value=(_repo_ref(), mock_adapter), + ): await _update_prd_proposal_pr( ticket_key="TEST-123", prd_content=unchanged_content, state=state, ) - mock_gh.create_or_update_file.assert_not_called() - mock_gh.create_issue_comment.assert_called_once() - comment_body = mock_gh.create_issue_comment.call_args[0][3] + mock_adapter.put_file.assert_not_called() + mock_adapter.create_comment.assert_called_once() + comment_body = mock_adapter.create_comment.call_args[0][2] assert "unchanged" in comment_body diff --git a/tests/unit/workflow/nodes/test_qa_handler.py b/tests/unit/workflow/nodes/test_qa_handler.py index ea4a1bf96..741ca2f1f 100644 --- a/tests/unit/workflow/nodes/test_qa_handler.py +++ b/tests/unit/workflow/nodes/test_qa_handler.py @@ -5,6 +5,7 @@ import pytest +from forge.integrations.source_control.contracts import Provider, RepositoryRef from forge.models.workflow import TicketType from forge.workflow.feature.state import create_initial_feature_state from forge.workflow.nodes.qa_handler import ( @@ -20,7 +21,9 @@ class TestExtractQuestionText: def test_strips_question_mark_prefix(self): """extract_question_text removes leading ? prefix.""" - assert extract_question_text("?What is this feature about?") == "What is this feature about?" + assert ( + extract_question_text("?What is this feature about?") == "What is this feature about?" + ) def test_strips_question_mark_prefix_with_whitespace(self): """extract_question_text handles ? with leading/trailing whitespace.""" @@ -434,12 +437,18 @@ async def test_closes_clients_on_error(self): @pytest.mark.asyncio async def test_posts_answer_to_github_pr_in_pr_mode(self): - """When prd_pr_number exists, Q&A answer goes to GitHub PR.""" + """When prd_pr_number exists, Q&A answer goes to GitHub PR via the adapter.""" mock_jira = create_mock_jira_client() mock_agent = create_mock_forge_agent() - mock_gh = MagicMock() - mock_gh.create_issue_comment = AsyncMock() - mock_gh.close = AsyncMock() + adapter = AsyncMock() + repo_ref = RepositoryRef( + id="org/proposals", + provider=Provider.GITHUB, + connection="c", + namespace="org/proposals", + default_branch="main", + change_request_mode="fork", + ) state = create_initial_feature_state( ticket_key="TEST-123", @@ -455,25 +464,31 @@ async def test_posts_answer_to_github_pr_in_pr_mode(self): with ( patch("forge.workflow.nodes.qa_handler.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.qa_handler.ForgeAgent", return_value=mock_agent), - patch("forge.workflow.nodes.qa_handler.GitHubClient", return_value=mock_gh), + patch("forge.workflow.nodes.qa_handler.get_adapter", return_value=(repo_ref, adapter)), ): await answer_question(state) - mock_gh.create_issue_comment.assert_called_once() - call_args = mock_gh.create_issue_comment.call_args[0] - assert call_args[0] == "org" - assert call_args[1] == "proposals" - assert call_args[2] == 7 + adapter.create_comment.assert_awaited_once() + call_args = adapter.create_comment.call_args[0] + assert call_args[0] is repo_ref + assert call_args[1].repository_id == "org/proposals" + assert call_args[1].native_id == 7 mock_jira.add_comment.assert_not_called() @pytest.mark.asyncio async def test_posts_spec_answer_to_github_pr_in_pr_mode(self): - """When spec_pr_number exists, spec Q&A answer goes to GitHub PR.""" + """When spec_pr_number exists, spec Q&A answer goes to GitHub PR via the adapter.""" mock_jira = create_mock_jira_client() mock_agent = create_mock_forge_agent() - mock_gh = MagicMock() - mock_gh.create_issue_comment = AsyncMock() - mock_gh.close = AsyncMock() + adapter = AsyncMock() + repo_ref = RepositoryRef( + id="org/proposals", + provider=Provider.GITHUB, + connection="c", + namespace="org/proposals", + default_branch="main", + change_request_mode="fork", + ) state = create_initial_feature_state( ticket_key="TEST-123", @@ -489,15 +504,15 @@ async def test_posts_spec_answer_to_github_pr_in_pr_mode(self): with ( patch("forge.workflow.nodes.qa_handler.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.qa_handler.ForgeAgent", return_value=mock_agent), - patch("forge.workflow.nodes.qa_handler.GitHubClient", return_value=mock_gh), + patch("forge.workflow.nodes.qa_handler.get_adapter", return_value=(repo_ref, adapter)), ): await answer_question(state) - mock_gh.create_issue_comment.assert_called_once() - call_args = mock_gh.create_issue_comment.call_args[0] - assert call_args[0] == "org" - assert call_args[1] == "proposals" - assert call_args[2] == 12 + adapter.create_comment.assert_awaited_once() + call_args = adapter.create_comment.call_args[0] + assert call_args[0] is repo_ref + assert call_args[1].repository_id == "org/proposals" + assert call_args[1].native_id == 12 mock_jira.add_comment.assert_not_called() @pytest.mark.asyncio @@ -599,7 +614,6 @@ def test_rca_includes_fix_options(self): assert "Option 2: Version cache entries" in content - class TestAnswerQuestionBugGates: """answer_question stays paused at all three new bug workflow gates.""" diff --git a/tests/unit/workflow/nodes/test_rebase.py b/tests/unit/workflow/nodes/test_rebase.py index d71428697..35b429700 100644 --- a/tests/unit/workflow/nodes/test_rebase.py +++ b/tests/unit/workflow/nodes/test_rebase.py @@ -6,7 +6,25 @@ import pytest -from forge.workflow.nodes.rebase import rebase_pr +from forge.integrations.source_control.contracts import ( + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + Provider, + RepositoryRef, +) +from forge.workflow.nodes.rebase import _fetch_pr_body, rebase_pr + + +def _repo_ref(): + return RepositoryRef( + id="acme/backend", + provider=Provider.GITHUB, + connection="c", + namespace="acme/backend", + default_branch="main", + change_request_mode="fork", + ) @pytest.mark.asyncio @@ -22,7 +40,7 @@ async def test_rebase_workspace_persists_identity_after_clone(tmp_path: Path) -> git = MagicMock() git.remote_branch_exists.return_value = False jira = MagicMock(close=AsyncMock()) - github = MagicMock(close=AsyncMock()) + adapter = AsyncMock() state = { "ticket_key": "TASK-123", "current_repo": "acme/backend", @@ -35,7 +53,7 @@ async def test_rebase_workspace_persists_identity_after_clone(tmp_path: Path) -> with ( patch("forge.workflow.nodes.rebase.get_settings"), patch("forge.workflow.nodes.rebase.JiraClient", return_value=jira), - patch("forge.workflow.nodes.rebase.GitHubClient", return_value=github), + patch("forge.workflow.nodes.rebase.get_adapter", return_value=(_repo_ref(), adapter)), patch("forge.workflow.nodes.rebase.get_workspace_manager", return_value=manager), patch("forge.workflow.nodes.rebase.GitOperations", return_value=git), ): @@ -45,3 +63,24 @@ async def test_rebase_workspace_persists_identity_after_clone(tmp_path: Path) -> '{"repo_name": "acme/backend", "ticket_key": "TASK-123"}\n' ) git.clone.assert_called_once_with() + + +@pytest.mark.asyncio +async def test_rebase_reads_body_via_adapter() -> None: + """_fetch_pr_body returns the PR body from the given adapter/identity.""" + adapter = AsyncMock() + adapter.get_change_request.return_value = ChangeRequest( + identity=ChangeRequestIdentity("c", "acme/widgets", 5), + url="u", + title="t", + body="PR body", + state=ChangeRequestState.OPEN, + source_branch="f", + target_branch="main", + ) + identity = ChangeRequestIdentity("c", "acme/widgets", 5) + + body = await _fetch_pr_body(adapter, _repo_ref(), identity) + + assert body == "PR body" + adapter.get_change_request.assert_awaited_once_with(_repo_ref(), identity) diff --git a/tests/unit/workflow/nodes/test_spec_pr.py b/tests/unit/workflow/nodes/test_spec_pr.py index 05efe5c10..dec35c67e 100644 --- a/tests/unit/workflow/nodes/test_spec_pr.py +++ b/tests/unit/workflow/nodes/test_spec_pr.py @@ -4,9 +4,39 @@ import pytest +from forge.integrations.source_control.contracts import ( + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + Provider, + RepositoryRef, + WriteTarget, +) +from forge.integrations.source_control.errors import ConflictError from forge.models.workflow import TicketType from forge.workflow.feature.state import create_initial_feature_state -from forge.integrations.github.client import PullRequestCreationResult + + +def _repo_ref(identifier: str = "org/proposals") -> RepositoryRef: + return RepositoryRef( + id=identifier, + provider=Provider.GITHUB, + connection="c", + namespace=identifier, + default_branch="main", + change_request_mode="fork", + ) + + +def _fork_ref(fork_owner: str, fork_repo: str) -> RepositoryRef: + return RepositoryRef( + id=f"{fork_owner}/{fork_repo}", + provider=Provider.GITHUB, + connection="c", + namespace=f"{fork_owner}/{fork_repo}", + default_branch="main", + change_request_mode="direct", + ) class TestCreateSpecProposalPr: @@ -14,29 +44,33 @@ class TestCreateSpecProposalPr: async def test_creates_branch_and_pr(self): from forge.workflow.nodes.spec_generation import _create_spec_proposal_pr - mock_gh = MagicMock() - mock_gh.get_or_create_fork = AsyncMock( - return_value={ - "owner": {"login": "forge-bot"}, - "name": "proposals", - "default_branch": "trunk", - } + mock_adapter = AsyncMock() + mock_adapter.resolve_default_branch = AsyncMock(return_value="trunk") + mock_adapter.ensure_write_target = AsyncMock( + return_value=WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="", + base_branch="main", + fork_owner="forge-bot", + fork_repo="proposals", + ) ) - mock_gh.get_repository = AsyncMock(return_value={"default_branch": "trunk"}) - mock_gh.sync_fork_with_upstream = AsyncMock(return_value=True) - mock_gh.get_file_contents = AsyncMock(return_value={"sha": "existing-sha"}) - mock_gh.create_branch = AsyncMock(return_value={"ref": "refs/heads/forge/spec/test-123"}) - mock_gh.create_or_update_file = AsyncMock(return_value={"content": {"sha": "filesha"}}) - mock_gh.create_pull_request = AsyncMock( - return_value=PullRequestCreationResult( - pr={ - "number": 12, - "html_url": "https://github.com/org/proposals/pull/12", - }, - created=True, + mock_adapter.create_branch = AsyncMock() + mock_adapter.put_file = AsyncMock() + mock_adapter.create_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="org/proposals", native_id=12 + ), + url="https://github.com/org/proposals/pull/12", + title="t", + body="b", + state=ChangeRequestState.OPEN, + source_branch="forge/spec/test-123", + target_branch="trunk", ) ) - mock_gh.close = AsyncMock() mock_jira = MagicMock() mock_jira.add_comment = AsyncMock() @@ -44,7 +78,10 @@ async def test_creates_branch_and_pr(self): mock_jira.close = AsyncMock() with ( - patch("forge.workflow.nodes.proposal_pr.GitHubClient", return_value=mock_gh), + patch( + "forge.workflow.nodes.proposal_pr.get_adapter", + return_value=(_repo_ref(), mock_adapter), + ), patch("forge.workflow.nodes.proposal_pr.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.proposal_pr.set_pr_ticket_index", @@ -63,24 +100,28 @@ async def test_creates_branch_and_pr(self): assert result["spec_pr_repo"] == "org/proposals" assert result["spec_pr_branch"] == "forge/spec/test-123" assert result["spec_pr_file_path"] == "TEST-123/design.md" - - mock_gh.create_branch.assert_called_once_with( - "forge-bot", "proposals", "forge/spec/test-123", base="trunk" - ) - mock_gh.sync_fork_with_upstream.assert_awaited_once_with( - "forge-bot", "proposals", branch="trunk" - ) assert result["spec_pr_fork_owner"] == "forge-bot" assert result["spec_pr_fork_repo"] == "proposals" - assert mock_gh.create_or_update_file.call_args.kwargs["sha"] == "existing-sha" - mock_gh.create_pull_request.assert_called_once() - pr_call_kwargs = mock_gh.create_pull_request.call_args[1] - assert pr_call_kwargs["owner"] == "org" - assert pr_call_kwargs["repo"] == "proposals" - assert pr_call_kwargs["head"] == "forge-bot:forge/spec/test-123" - assert pr_call_kwargs["base"] == "trunk" - assert "# My Spec" not in pr_call_kwargs["body"] - assert "TEST-123/design.md" in pr_call_kwargs["body"] + + mock_adapter.create_branch.assert_called_once_with( + _fork_ref("forge-bot", "proposals"), "forge/spec/test-123", "trunk" + ) + mock_adapter.put_file.assert_called_once() + put_args = mock_adapter.put_file.call_args[0] + assert put_args[0] == _fork_ref("forge-bot", "proposals") + assert put_args[1] == "TEST-123/design.md" + assert put_args[2] == "# My Spec" + assert put_args[4] == "forge/spec/test-123" + + mock_adapter.create_change_request.assert_called_once() + cr_args, cr_kwargs = mock_adapter.create_change_request.call_args + assert cr_args[0] == _repo_ref() + assert cr_args[1].head_ref == "forge/spec/test-123" + assert cr_args[1].base_branch == "trunk" + assert cr_args[1].fork_owner == "forge-bot" + assert "# My Spec" not in cr_kwargs["body"] + assert "TEST-123/design.md" in cr_kwargs["body"] + mock_jira.add_comment.assert_called_once() mock_jira.set_workflow_label.assert_called_once() mock_index.assert_called_once() @@ -89,26 +130,33 @@ async def test_creates_branch_and_pr(self): async def test_creates_pr_with_custom_path(self): from forge.workflow.nodes.spec_generation import _create_spec_proposal_pr - mock_gh = MagicMock() - mock_gh.get_or_create_fork = AsyncMock( - return_value={"owner": {"login": "forge-bot"}, "name": "proposals"} + mock_adapter = AsyncMock() + mock_adapter.resolve_default_branch = AsyncMock(return_value="main") + mock_adapter.ensure_write_target = AsyncMock( + return_value=WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="", + base_branch="main", + fork_owner="forge-bot", + fork_repo="proposals", + ) ) - # Omitting a custom default branch exercises the upstream "main" fallback. - mock_gh.get_repository = AsyncMock(return_value={}) - mock_gh.sync_fork_with_upstream = AsyncMock(return_value=True) - mock_gh.get_file_contents = AsyncMock(return_value=None) - mock_gh.create_branch = AsyncMock(return_value={"ref": "refs/heads/forge/spec/test-456"}) - mock_gh.create_or_update_file = AsyncMock(return_value={"content": {"sha": "filesha"}}) - mock_gh.create_pull_request = AsyncMock( - return_value=PullRequestCreationResult( - pr={ - "number": 15, - "html_url": "https://github.com/org/proposals/pull/15", - }, - created=True, + mock_adapter.create_branch = AsyncMock() + mock_adapter.put_file = AsyncMock() + mock_adapter.create_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity( + connection="c", repository_id="org/proposals", native_id=15 + ), + url="https://github.com/org/proposals/pull/15", + title="t", + body="b", + state=ChangeRequestState.OPEN, + source_branch="forge/spec/test-456", + target_branch="main", ) ) - mock_gh.close = AsyncMock() mock_jira = MagicMock() mock_jira.add_comment = AsyncMock() @@ -116,7 +164,10 @@ async def test_creates_pr_with_custom_path(self): mock_jira.close = AsyncMock() with ( - patch("forge.workflow.nodes.proposal_pr.GitHubClient", return_value=mock_gh), + patch( + "forge.workflow.nodes.proposal_pr.get_adapter", + return_value=(_repo_ref(), mock_adapter), + ), patch("forge.workflow.nodes.proposal_pr.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.proposal_pr.set_pr_ticket_index", @@ -132,28 +183,28 @@ async def test_creates_pr_with_custom_path(self): ) assert result["spec_pr_file_path"] == "enhancements/TEST-456/design.md" - pr_call_kwargs = mock_gh.create_pull_request.call_args[1] - assert "enhancements/TEST-456/design.md" in pr_call_kwargs["body"] - assert pr_call_kwargs["base"] == "main" - assert mock_gh.create_or_update_file.call_args.kwargs["sha"] is None + cr_args, cr_kwargs = mock_adapter.create_change_request.call_args + assert "enhancements/TEST-456/design.md" in cr_kwargs["body"] + assert cr_args[1].base_branch == "main" @pytest.mark.asyncio async def test_stops_when_fork_cannot_be_synchronized(self): from forge.workflow.nodes.spec_generation import _create_spec_proposal_pr - mock_gh = MagicMock() - mock_gh.get_or_create_fork = AsyncMock( - return_value={"owner": {"login": "forge-bot"}, "name": "proposals"} + mock_adapter = AsyncMock() + mock_adapter.resolve_default_branch = AsyncMock(return_value="main") + mock_adapter.ensure_write_target = AsyncMock( + side_effect=ConflictError("Could not synchronize proposal fork") ) - mock_gh.get_repository = AsyncMock(return_value={"default_branch": "main"}) - mock_gh.sync_fork_with_upstream = AsyncMock(return_value=False) - mock_gh.close = AsyncMock() mock_jira = MagicMock(close=AsyncMock()) with ( - patch("forge.workflow.nodes.proposal_pr.GitHubClient", return_value=mock_gh), + patch( + "forge.workflow.nodes.proposal_pr.get_adapter", + return_value=(_repo_ref(), mock_adapter), + ), patch("forge.workflow.nodes.proposal_pr.JiraClient", return_value=mock_jira), - pytest.raises(RuntimeError, match="Could not synchronize proposal fork"), + pytest.raises(ConflictError, match="Could not synchronize proposal fork"), ): await _create_spec_proposal_pr( ticket_key="TEST-FAIL", @@ -162,9 +213,9 @@ async def test_stops_when_fork_cannot_be_synchronized(self): proposals_repo="org/proposals", ) - mock_gh.create_branch.assert_not_called() - mock_gh.create_or_update_file.assert_not_called() - mock_gh.create_pull_request.assert_not_called() + mock_adapter.create_branch.assert_not_called() + mock_adapter.put_file.assert_not_called() + mock_adapter.create_change_request.assert_not_called() class TestUpdateSpecProposalPr: @@ -172,13 +223,10 @@ class TestUpdateSpecProposalPr: async def test_updates_file_on_branch(self): from forge.workflow.nodes.spec_generation import _update_spec_proposal_pr - mock_gh = MagicMock() - mock_gh.get_file_contents = AsyncMock( - return_value={"sha": "oldsha", "path": "TEST-123/design.md"} - ) - mock_gh.create_or_update_file = AsyncMock(return_value={"content": {"sha": "newsha"}}) - mock_gh.create_issue_comment = AsyncMock() - mock_gh.close = AsyncMock() + mock_adapter = AsyncMock() + mock_adapter.get_file = AsyncMock(return_value="# Old Spec") + mock_adapter.put_file = AsyncMock() + mock_adapter.create_comment = AsyncMock() state = create_initial_feature_state( ticket_key="TEST-123", @@ -192,36 +240,42 @@ async def test_updates_file_on_branch(self): spec_pr_file_path="TEST-123/design.md", ) - with patch("forge.workflow.nodes.proposal_pr.GitHubClient", return_value=mock_gh): + with patch( + "forge.workflow.nodes.proposal_pr.get_adapter", + return_value=(_repo_ref(), mock_adapter), + ): await _update_spec_proposal_pr( ticket_key="TEST-123", spec_content="# Revised Spec", state=state, ) - mock_gh.get_file_contents.assert_called_once_with( - "forge-bot", "proposals", "TEST-123/design.md", "forge/spec/test-123" + mock_adapter.get_file.assert_called_once_with( + _fork_ref("forge-bot", "proposals"), "TEST-123/design.md", "forge/spec/test-123" ) - mock_gh.create_or_update_file.assert_called_once() - call_kwargs = mock_gh.create_or_update_file.call_args[1] - assert call_kwargs["sha"] == "oldsha" - assert call_kwargs["path"] == "TEST-123/design.md" - mock_gh.create_issue_comment.assert_called_once_with( - "org", - "proposals", - 12, - "Specification has been revised based on feedback. Please review the updated version.", + mock_adapter.put_file.assert_called_once() + put_args = mock_adapter.put_file.call_args[0] + assert put_args[0] == _fork_ref("forge-bot", "proposals") + assert put_args[1] == "TEST-123/design.md" + assert put_args[2] == "# Revised Spec" + mock_adapter.create_comment.assert_called_once() + comment_args = mock_adapter.create_comment.call_args[0] + assert comment_args[0] == _repo_ref() + assert comment_args[1].native_id == 12 + assert ( + comment_args[2] + == "Specification has been revised based on feedback. Please review the updated version." ) @pytest.mark.asyncio async def test_updates_legacy_upstream_branch_without_fork_state(self): from forge.workflow.nodes.spec_generation import _update_spec_proposal_pr - mock_gh = MagicMock() - mock_gh.get_file_contents = AsyncMock(return_value={"sha": "oldsha"}) - mock_gh.create_or_update_file = AsyncMock() - mock_gh.create_issue_comment = AsyncMock() - mock_gh.close = AsyncMock() + mock_adapter = AsyncMock() + mock_adapter.get_file = AsyncMock(return_value="# Old") + mock_adapter.put_file = AsyncMock() + mock_adapter.create_comment = AsyncMock() + state = create_initial_feature_state( ticket_key="TEST-LEGACY", ticket_type=TicketType.FEATURE, @@ -231,11 +285,14 @@ async def test_updates_legacy_upstream_branch_without_fork_state(self): spec_pr_file_path="TEST-LEGACY/design.md", ) - with patch("forge.workflow.nodes.proposal_pr.GitHubClient", return_value=mock_gh): + with patch( + "forge.workflow.nodes.proposal_pr.get_adapter", + return_value=(_repo_ref(), mock_adapter), + ): await _update_spec_proposal_pr("TEST-LEGACY", "# Revised", state) - mock_gh.get_file_contents.assert_awaited_once_with( - "org", "proposals", "TEST-LEGACY/design.md", "forge/spec/test-legacy" + mock_adapter.get_file.assert_awaited_once_with( + _repo_ref(), "TEST-LEGACY/design.md", "forge/spec/test-legacy" ) - assert mock_gh.create_or_update_file.call_args.kwargs["owner"] == "org" - assert mock_gh.create_or_update_file.call_args.kwargs["repo"] == "proposals" + put_args = mock_adapter.put_file.call_args[0] + assert put_args[0] == _repo_ref() diff --git a/tests/unit/workflow/nodes/test_task_takeover_execution.py b/tests/unit/workflow/nodes/test_task_takeover_execution.py index 80a641983..ed7951e5e 100644 --- a/tests/unit/workflow/nodes/test_task_takeover_execution.py +++ b/tests/unit/workflow/nodes/test_task_takeover_execution.py @@ -29,6 +29,8 @@ def _make_state( "plan_content": plan_content, "implemented_tasks": implemented_tasks or [], "context": {"branch_name": "forge/TASK-123", "guardrails": ""}, + "fork_owner": "forge-bot", + "fork_repo": "backend", } @@ -117,8 +119,9 @@ async def test_successful_execution(self) -> None: assert "Approved Implementation Plan" in kwargs["task_description"] assert "inject at least one new or modified test file" in kwargs["task_description"] assert "Current repository: `acme/backend`" in kwargs["task_description"] - assert "Do not search for, create, or modify files assigned to other repositories" in ( - kwargs["task_description"] + assert ( + "Do not search for, create, or modify files assigned to other repositories" + in (kwargs["task_description"]) ) assert "config" not in kwargs diff --git a/tests/unit/workflow/nodes/test_workspace_setup.py b/tests/unit/workflow/nodes/test_workspace_setup.py index 38213b4fd..316052ab0 100644 --- a/tests/unit/workflow/nodes/test_workspace_setup.py +++ b/tests/unit/workflow/nodes/test_workspace_setup.py @@ -9,10 +9,16 @@ import httpx import pytest +from forge.integrations.source_control.contracts import ( + GitCredentials, + Provider, + RepositoryRef, + WriteTarget, +) +from forge.integrations.source_control.errors import SourceControlError from forge.models.workflow import ForgeLabel from forge.workflow.feature.state import create_initial_feature_state from forge.workflow.nodes.workspace_setup import prepare_workspace, setup_workspace -from forge.workspace.manager import WorkspaceManager def create_mock_jira_client(): @@ -67,18 +73,41 @@ def create_mock_guardrails_loader(): return mock +def _repo_ref(identifier: str) -> RepositoryRef: + return RepositoryRef( + id=identifier, + provider=Provider.GITHUB, + connection="c", + namespace=identifier, + default_branch="main", + change_request_mode="fork", + ) + + @pytest.fixture(autouse=True) -def mock_workspace_github(): - """Keep workspace tests isolated from GitHub fork operations.""" - github = MagicMock() - github.get_repository = AsyncMock(return_value={"default_branch": "main"}) - github.get_or_create_fork = AsyncMock( - return_value={"owner": {"login": "fork-owner"}, "name": "test-repo"} +def mock_workspace_adapter(): + """Keep workspace tests isolated from source-control adapter calls.""" + adapter = MagicMock() + adapter.resolve_default_branch = AsyncMock(return_value="main") + adapter.ensure_write_target = AsyncMock( + return_value=WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="", + base_branch="main", + fork_owner="fork-owner", + fork_repo="test-repo", + ) ) - github.sync_fork_with_upstream = AsyncMock(return_value=True) - github.close = AsyncMock() - with patch("forge.workflow.nodes.workspace_setup.GitHubClient", return_value=github): - yield github + adapter.get_git_credentials = AsyncMock( + return_value=GitCredentials(host="github.com", token="test-token") + ) + + def _get_adapter(identifier): + return _repo_ref(identifier), adapter + + with patch("forge.workflow.nodes.workspace_setup.get_adapter", side_effect=_get_adapter): + yield adapter class TestWorkspaceSetupStatusComment: @@ -253,47 +282,6 @@ async def test_workspace_setup_transitions_tasks(self): mock_jira.transition_issue.assert_any_call("AISOS-101", "In Progress") mock_jira.transition_issue.assert_any_call("AISOS-102", "In Progress") - @pytest.mark.asyncio - async def test_workspace_setup_transitions_epics_and_parent_epic(self): - """Should transition tasks, epic keys, and parent epic to In Progress.""" - mock_jira = create_mock_jira_client() - mock_issue = MagicMock() - mock_issue.parent_key = "EPIC-999" - mock_jira.get_issue = AsyncMock(return_value=mock_issue) - - mock_manager, mock_workspace = create_mock_workspace_manager() - mock_git = create_mock_git_operations() - mock_guardrails_loader = create_mock_guardrails_loader() - - state = create_initial_feature_state( - ticket_key="TEST-101", - current_repo="owner/test-repo", - task_keys=["TASK-1", "TASK-2"], - ) - state["epic_keys"] = ["EPIC-1", "EPIC-2"] - - with ( - patch("forge.workflow.nodes.workspace_setup.JiraClient", return_value=mock_jira), - patch( - "forge.workflow.nodes.workspace_setup.get_workspace_manager", - return_value=mock_manager, - ), - patch("forge.workflow.nodes.workspace_setup.GitOperations", return_value=mock_git), - patch("forge.workflow.nodes.workspace_setup.GuardrailsLoader", mock_guardrails_loader), - ): - await setup_workspace(state) - - # Asserts that transition_issue is called for TASK-1, TASK-2, EPIC-1, EPIC-2, and EPIC-999 - assert mock_jira.transition_issue.call_count == 5 - mock_jira.transition_issue.assert_any_call("TASK-1", "In Progress") - mock_jira.transition_issue.assert_any_call("TASK-2", "In Progress") - mock_jira.transition_issue.assert_any_call("EPIC-1", "In Progress") - mock_jira.transition_issue.assert_any_call("EPIC-2", "In Progress") - mock_jira.transition_issue.assert_any_call("EPIC-999", "In Progress") - - # Asserts that get_issue was called with the Feature's ticket key - mock_jira.get_issue.assert_called_once_with("TEST-101") - class TestWorkspaceSetupErrorHandling: """Test cases for workspace setup error handling.""" @@ -337,9 +325,11 @@ async def test_workspace_setup_continues_on_jira_failure(self, caplog): mock_jira.close.assert_called_once() @pytest.mark.asyncio - async def test_workspace_setup_fails_when_fork_cannot_be_created(self, mock_workspace_github): + async def test_workspace_setup_fails_when_fork_cannot_be_created(self, mock_workspace_adapter): """Implementation must not start without its durable backup remote.""" - mock_workspace_github.get_or_create_fork.side_effect = RuntimeError("fork creation denied") + mock_workspace_adapter.ensure_write_target.side_effect = SourceControlError( + "fork creation denied" + ) mock_jira = create_mock_jira_client() mock_manager, _ = create_mock_workspace_manager() mock_git = create_mock_git_operations() @@ -369,7 +359,7 @@ class TestWorkspaceSetupForkBootstrap: """Tests for creating and checkpointing the implementation backup fork.""" @pytest.mark.asyncio - async def test_creates_fork_remote_before_implementation(self, mock_workspace_github): + async def test_creates_fork_remote_before_implementation(self, mock_workspace_adapter): mock_jira = create_mock_jira_client() mock_manager, _ = create_mock_workspace_manager() mock_git = create_mock_git_operations() @@ -390,9 +380,11 @@ async def test_creates_fork_remote_before_implementation(self, mock_workspace_gi ): result = await setup_workspace(state) - mock_workspace_github.get_or_create_fork.assert_awaited_once_with("upstream", "repo") - mock_workspace_github.sync_fork_with_upstream.assert_awaited_once_with( - "fork-owner", "test-repo", branch="main" + mock_workspace_adapter.resolve_default_branch.assert_awaited_once_with( + _repo_ref("upstream/repo") + ) + mock_workspace_adapter.ensure_write_target.assert_awaited_once_with( + _repo_ref("upstream/repo") ) mock_git.add_fork_remote.assert_called_once_with("fork-owner", "test-repo") mock_git.push_to_fork.assert_called_once_with() @@ -402,7 +394,7 @@ async def test_creates_fork_remote_before_implementation(self, mock_workspace_gi @pytest.mark.asyncio async def test_initial_branch_push_failure_prevents_implementation_handoff( - self, mock_workspace_github + self, mock_workspace_adapter ): mock_jira = create_mock_jira_client() mock_manager, _ = create_mock_workspace_manager() @@ -425,14 +417,16 @@ async def test_initial_branch_push_failure_prevents_implementation_handoff( ): result = await setup_workspace(state) - mock_workspace_github.get_or_create_fork.assert_awaited_once_with("upstream", "repo") + mock_workspace_adapter.ensure_write_target.assert_awaited_once_with( + _repo_ref("upstream/repo") + ) mock_git.push_to_fork.assert_called_once_with() assert result["current_node"] == "setup_workspace" assert result["retry_count"] == 1 assert "invalid refspec" in result["last_error"] @pytest.mark.asyncio - async def test_existing_fork_branch_is_checked_out_without_push(self, mock_workspace_github): + async def test_existing_fork_branch_is_checked_out_without_push(self, mock_workspace_adapter): mock_jira = create_mock_jira_client() mock_manager, mock_workspace = create_mock_workspace_manager() mock_git = create_mock_git_operations() @@ -454,8 +448,8 @@ async def test_existing_fork_branch_is_checked_out_without_push(self, mock_works ): result = await setup_workspace(state) - mock_workspace_github.sync_fork_with_upstream.assert_awaited_once_with( - "fork-owner", "test-repo", branch="main" + mock_workspace_adapter.ensure_write_target.assert_awaited_once_with( + _repo_ref("upstream/repo") ) mock_git.remote_branch_exists.assert_called_once_with( mock_workspace.branch_name, remote="fork" @@ -465,11 +459,107 @@ async def test_existing_fork_branch_is_checked_out_without_push(self, mock_works mock_git.push_to_fork.assert_not_called() assert result["current_node"] == "implementation" + @pytest.mark.asyncio + async def test_direct_mode_pushes_to_origin_without_fork_remote(self, mock_workspace_adapter): + """change_request_mode == "direct" has no fork identity — setup must not + build a fork remote from empty owner/repo, and must push to origin.""" + mock_workspace_adapter.ensure_write_target = AsyncMock( + return_value=WriteTarget( + clone_url="", + push_remote_name="origin", + head_ref="", + base_branch="main", + fork_owner=None, + fork_repo=None, + ) + ) + mock_jira = create_mock_jira_client() + mock_manager, mock_workspace = create_mock_workspace_manager() + mock_git = create_mock_git_operations() + mock_guardrails_loader = create_mock_guardrails_loader() + state = create_initial_feature_state( + ticket_key="TEST-127", + current_repo="upstream/repo", + ) + + with ( + patch("forge.workflow.nodes.workspace_setup.JiraClient", return_value=mock_jira), + patch( + "forge.workflow.nodes.workspace_setup.get_workspace_manager", + return_value=mock_manager, + ), + patch("forge.workflow.nodes.workspace_setup.GitOperations", return_value=mock_git), + patch("forge.workflow.nodes.workspace_setup.GuardrailsLoader", mock_guardrails_loader), + ): + result = await setup_workspace(state) + + mock_git.add_fork_remote.assert_not_called() + mock_git.remote_branch_exists.assert_called_once_with( + mock_workspace.branch_name, remote="origin" + ) + mock_git.push_to_fork.assert_not_called() + mock_git.push.assert_called_once_with(force=False, check_conflicts=False) + assert result["fork_owner"] == "" + assert result["fork_repo"] == "" + assert result["current_node"] == "implementation" + + @pytest.mark.asyncio + async def test_repos_yaml_alias_is_canonicalized_before_cloning(self, mock_workspace_adapter): + """A repos.yaml alias (e.g. "payments-api") has no "/" and must be + resolved to its canonical "owner/repo" namespace before cloning, so + the PR state this run saves (keyed by current_repo) later matches + webhook lookups (keyed by event.repo_ref.namespace).""" + alias_repo_ref = RepositoryRef( + id="payments-api", + provider=Provider.GITHUB, + connection="c", + namespace="acme/payments", + default_branch="main", + change_request_mode="fork", + ) + mock_jira = create_mock_jira_client() + mock_manager, _ = create_mock_workspace_manager() + mock_git = create_mock_git_operations() + mock_guardrails_loader = create_mock_guardrails_loader() + state = create_initial_feature_state( + ticket_key="TEST-ALIAS", + current_repo="payments-api", + ) + state["tasks_by_repo"] = {"payments-api": ["TASK-1", "TASK-2"]} + state["repos_to_process"] = ["payments-api"] + + with ( + patch("forge.workflow.nodes.workspace_setup.JiraClient", return_value=mock_jira), + patch( + "forge.workflow.nodes.workspace_setup.get_workspace_manager", + return_value=mock_manager, + ), + patch("forge.workflow.nodes.workspace_setup.GitOperations", return_value=mock_git), + patch("forge.workflow.nodes.workspace_setup.GuardrailsLoader", mock_guardrails_loader), + patch( + "forge.workflow.nodes.workspace_setup.get_adapter", + return_value=(alias_repo_ref, mock_workspace_adapter), + ), + ): + result = await setup_workspace(state) + + mock_manager.create_workspace.assert_called_once_with( + repo_name="acme/payments", ticket_key="TEST-ALIAS" + ) + assert result["current_repo"] == "acme/payments" + assert result["current_node"] == "implementation" + # tasks_by_repo/repos_to_process must move to the canonical key in + # lockstep with current_repo, or implementation's task lookup and + # route_after_pr's completion matching desync from the alias. + assert result["tasks_by_repo"] == {"acme/payments": ["TASK-1", "TASK-2"]} + assert result["repos_to_process"] == ["acme/payments"] + class TestPrepareWorkspaceRecovery: """Tests for prepare_workspace workspace sync/recreation behavior.""" - def test_sync_failure_recreates_workspace_from_fork(self, tmp_path): + @pytest.mark.asyncio + async def test_sync_failure_recreates_workspace_from_fork(self, tmp_path): """A workspace that cannot sync is deleted and cloned fresh from the fork.""" workspace_path = tmp_path / "forge-TEST-123-org-repo" workspace_path.mkdir() @@ -483,13 +573,6 @@ def test_sync_failure_recreates_workspace_from_fork(self, tmp_path): fork_owner="forge-bot", fork_repo="repo", context={"branch_name": "forge/test-123"}, - handoffs={ - "org/repo": { - "content": "prior task context", - "task_key": "TEST-122", - "captured_at": "2026-08-06T00:00:00+00:00", - } - }, ) old_git = MagicMock() @@ -504,7 +587,7 @@ def test_sync_failure_recreates_workspace_from_fork(self, tmp_path): side_effect=[old_git, new_git], ), ): - result_path, result_git = prepare_workspace(state) + result_path, result_git = await prepare_workspace(state) assert result_path == str(workspace_path) assert result_git is new_git @@ -514,23 +597,32 @@ def test_sync_failure_recreates_workspace_from_fork(self, tmp_path): new_git.add_fork_remote.assert_called_once_with("forge-bot", "repo") new_git.checkout_branch.assert_called_once_with("forge/test-123", remote="fork") assert new_git.workspace_recreated is True - assert (workspace_path / ".forge" / "handoff.md").read_text() == "prior task context" - def test_recovery_uses_container_aware_cleanup_for_old_workspace(self, tmp_path): - workspace_path = tmp_path / "forge-TEST-125-org-repo" + @pytest.mark.asyncio + async def test_sync_failure_recreates_direct_mode_workspace_from_origin(self, tmp_path): + """Direct-mode recovery (no fork identity) clones and checks out from origin. + + With no fork_owner/fork_repo, recreation must not build a fork remote + from empty owner/repo; it checks the branch out directly from origin, + relying on checkout_branch to fetch the ref that a single-branch clone + would otherwise miss. + """ + workspace_path = tmp_path / "forge-TEST-127-org-repo" workspace_path.mkdir() + stale_file = workspace_path / "stale.txt" + stale_file.write_text("dirty") state = create_initial_feature_state( - ticket_key="TEST-125", + ticket_key="TEST-127", current_repo="org/repo", workspace_path=str(workspace_path), - fork_owner="forge-bot", - fork_repo="repo", - context={"branch_name": "forge/test-125"}, + fork_owner=None, + fork_repo=None, + context={"branch_name": "forge/test-direct"}, ) old_git = MagicMock() - old_git.pull_rebase.side_effect = RuntimeError("sync failed") + old_git.pull_rebase.side_effect = RuntimeError("workspace sync failure") new_git = MagicMock() settings = MagicMock(workspace_base_dir=str(tmp_path)) @@ -540,15 +632,20 @@ def test_recovery_uses_container_aware_cleanup_for_old_workspace(self, tmp_path) "forge.workflow.nodes.workspace_setup.GitOperations", side_effect=[old_git, new_git], ), - patch.object(WorkspaceManager, "remove_path") as remove_path, ): - result_path, _ = prepare_workspace(state) + result_path, result_git = await prepare_workspace(state) assert result_path == str(workspace_path) - remove_path.assert_called_once() - assert "-old-" in remove_path.call_args.args[0].name + assert result_git is new_git + assert not stale_file.exists() + old_git.pull_rebase.assert_called_once_with(remote="origin") + new_git.clone.assert_called_once() + new_git.add_fork_remote.assert_not_called() + new_git.checkout_branch.assert_called_once_with("forge/test-direct", remote="origin") + assert new_git.workspace_recreated is True - def test_failed_replacement_preserves_existing_workspace(self, tmp_path): + @pytest.mark.asyncio + async def test_failed_replacement_preserves_existing_workspace(self, tmp_path): """A failed recovery clone must not delete the only local commit.""" workspace_path = tmp_path / "forge-TEST-124-org-repo" workspace_path.mkdir() @@ -578,11 +675,12 @@ def test_failed_replacement_preserves_existing_workspace(self, tmp_path): ), pytest.raises(RuntimeError, match="clone failed"), ): - prepare_workspace(state) + await prepare_workspace(state) assert local_commit.read_text() == "not pushed yet" - def test_backup_cleanup_retries_directory_not_empty(self, tmp_path): + @pytest.mark.asyncio + async def test_backup_cleanup_retries_directory_not_empty(self, tmp_path): """A transient ENOTEMPTY race during backup deletion is retried.""" workspace_path = tmp_path / "forge-TEST-125-org-repo" workspace_path.mkdir() @@ -602,13 +700,13 @@ def test_backup_cleanup_retries_directory_not_empty(self, tmp_path): real_rmtree = shutil.rmtree cleanup_calls = 0 - def transient_remove(path): + def transient_rmtree(path, *args, **kwargs): nonlocal cleanup_calls if Path(path).name.startswith(f".{workspace_path.name}-old-"): cleanup_calls += 1 if cleanup_calls == 1: raise OSError(errno.ENOTEMPTY, "Directory not empty", path) - return real_rmtree(path) + return real_rmtree(path, *args, **kwargs) with ( patch("forge.workflow.nodes.workspace_setup.get_settings", return_value=settings), @@ -616,17 +714,21 @@ def transient_remove(path): "forge.workflow.nodes.workspace_setup.GitOperations", side_effect=[old_git, new_git], ), - patch.object(WorkspaceManager, "remove_path", side_effect=transient_remove), + patch( + "forge.workflow.nodes.workspace_setup.WorkspaceManager.remove_path", + side_effect=transient_rmtree, + ), patch("forge.workflow.nodes.workspace_setup.time.sleep") as sleep, ): - result_path, result_git = prepare_workspace(state) + result_path, result_git = await prepare_workspace(state) assert result_path == str(workspace_path) assert result_git is new_git assert cleanup_calls == 2 sleep.assert_called_once() - def test_backup_cleanup_failure_does_not_fail_recovery(self, tmp_path, caplog): + @pytest.mark.asyncio + async def test_backup_cleanup_failure_does_not_fail_recovery(self, tmp_path, caplog): """Post-swap cleanup errors do not mask successful workspace recovery.""" workspace_path = tmp_path / "forge-TEST-126-org-repo" workspace_path.mkdir() @@ -645,10 +747,10 @@ def test_backup_cleanup_failure_does_not_fail_recovery(self, tmp_path, caplog): settings = MagicMock(workspace_base_dir=str(tmp_path)) real_rmtree = shutil.rmtree - def persistent_remove(path): + def persistent_rmtree(path, *args, **kwargs): if Path(path).name.startswith(f".{workspace_path.name}-old-"): raise OSError(errno.ENOTEMPTY, "Directory not empty", path) - return real_rmtree(path) + return real_rmtree(path, *args, **kwargs) with ( patch("forge.workflow.nodes.workspace_setup.get_settings", return_value=settings), @@ -656,11 +758,14 @@ def persistent_remove(path): "forge.workflow.nodes.workspace_setup.GitOperations", side_effect=[old_git, new_git], ), - patch.object(WorkspaceManager, "remove_path", side_effect=persistent_remove), + patch( + "forge.workflow.nodes.workspace_setup.WorkspaceManager.remove_path", + side_effect=persistent_rmtree, + ), patch("forge.workflow.nodes.workspace_setup.time.sleep"), caplog.at_level("WARNING"), ): - result_path, result_git = prepare_workspace(state) + result_path, result_git = await prepare_workspace(state) assert result_path == str(workspace_path) assert result_git is new_git diff --git a/tests/unit/workflow/test_base.py b/tests/unit/workflow/test_base.py index 9982b3cd0..9daf96b4a 100644 --- a/tests/unit/workflow/test_base.py +++ b/tests/unit/workflow/test_base.py @@ -63,7 +63,6 @@ def test_pr_state_has_required_fields(self): assert "repos_completed" in hints assert "implemented_tasks" in hints assert "current_task_key" in hints - assert "handoffs" in hints class TestCIIntegrationState: @@ -148,7 +147,7 @@ def state_schema(self): from forge.workflow.base import BaseState return BaseState - def matches(self, ticket_type, labels, event): + def matches(self, _ticket_type, _labels, _event): return True def build_graph(self): @@ -176,9 +175,12 @@ def build_graph(self): class TestSharedResumeMap: """Tests for resolve_shared_resume_node.""" - def test_wait_for_ci_gate_is_not_mapped(self): - """wait_for_ci_gate was removed — it must not resolve to any node.""" + def test_wait_for_ci_gate_resolves_to_human_review_gate(self): + """wait_for_ci_gate was removed and merged into human_review_gate. A + compatibility alias must still resolve it so a ticket checkpointed at + wait_for_ci_gate before the merge resumes instead of silently + restarting its whole workflow.""" from forge.workflow.utils import resolve_shared_resume_node result = resolve_shared_resume_node("wait_for_ci_gate") - assert result is None + assert result == "human_review_gate" diff --git a/tests/unit/workflow/test_ci_gate_skip.py b/tests/unit/workflow/test_ci_gate_skip.py index dcae4c50f..fcb02f0e1 100644 --- a/tests/unit/workflow/test_ci_gate_skip.py +++ b/tests/unit/workflow/test_ci_gate_skip.py @@ -1,51 +1,80 @@ """Tests for CI gate skip via GitHub PR comment (proposal 005).""" +from datetime import UTC, datetime from unittest.mock import AsyncMock, MagicMock, patch import pytest -from tests.fixtures.workflow_states import make_workflow_state +from forge.integrations.source_control.contracts import ( + Actor, + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + CheckConclusion, + CheckRun, + CheckStatus, + EventKind, + NormalizedEvent, + Provider, + RepositoryRef, + ReviewComment, +) from forge.models.events import EventSource from forge.orchestrator.worker import OrchestratorWorker -from forge.queue.models import QueueMessage +from forge.queue.models import QueueMessage, normalized_event_to_dict +from tests.fixtures.workflow_states import make_workflow_state # ── Helpers ─────────────────────────────────────────────────────────────────── -def _skip_gate_message(base: QueueMessage, check_name: str) -> QueueMessage: - """GitHub issue_comment event with /forge skip-gate command.""" +def _comment_message(base: QueueMessage, body: str) -> QueueMessage: + """Source-control comment event (COMMENT_CREATED) carrying a PR comment body.""" + event = NormalizedEvent( + id=base.event_id, + kind=EventKind.COMMENT_CREATED, + repo_ref=RepositoryRef( + id="org/repo", + provider=Provider.GITHUB, + connection="default-github", + namespace="org/repo", + default_branch="main", + change_request_mode="fork", + ), + actor=Actor(login="eshulman2", is_bot=False), + received_at=datetime(2026, 1, 1, tzinfo=UTC), + change_request=ChangeRequest( + identity=ChangeRequestIdentity( + connection="default-github", repository_id="org/repo", native_id=42 + ), + url="https://github.com/org/repo/pull/42", + title="t", + body="", + state=ChangeRequestState.OPEN, + source_branch="feature", + target_branch="main", + draft=False, + ), + comment=ReviewComment(id="c1", body=body, author="eshulman2"), + ) return QueueMessage( message_id=base.message_id, event_id=base.event_id, - source=EventSource.GITHUB, - event_type="issue_comment:created", # GitHub appends :action + source=EventSource.SOURCE_CONTROL, + event_type="comment_created", ticket_key=base.ticket_key, - payload={ - **base.payload, - "comment": {"body": f"/forge skip-gate {check_name}"}, - "issue": {"number": 42, "pull_request": {}}, - "repository": {"full_name": "org/repo"}, - "sender": {"login": "eshulman2"}, - }, + payload={}, + normalized_event=normalized_event_to_dict(event), ) +def _skip_gate_message(base: QueueMessage, check_name: str) -> QueueMessage: + """Source-control comment event with /forge skip-gate command.""" + return _comment_message(base, f"/forge skip-gate {check_name}") + + def _unskip_gate_message(base: QueueMessage, check_name: str) -> QueueMessage: - """GitHub issue_comment event with /forge unskip-gate command.""" - return QueueMessage( - message_id=base.message_id, - event_id=base.event_id, - source=EventSource.GITHUB, - event_type="issue_comment:created", - ticket_key=base.ticket_key, - payload={ - **base.payload, - "comment": {"body": f"/forge unskip-gate {check_name}"}, - "issue": {"number": 42, "pull_request": {}}, - "repository": {"full_name": "org/repo"}, - "sender": {"login": "eshulman2"}, - }, - ) + """Source-control comment event with /forge unskip-gate command.""" + return _comment_message(base, f"/forge unskip-gate {check_name}") @pytest.fixture @@ -58,12 +87,10 @@ def base_message(): return QueueMessage( message_id="1234567890-0", event_id="test-event-001", - source=EventSource.GITHUB, - event_type="issue_comment", + source=EventSource.SOURCE_CONTROL, + event_type="comment_created", ticket_key="TEST-123", - payload={ - "issue": {"key": "TEST-123", "fields": {"issuetype": {"name": "Feature"}}}, - }, + payload={}, ) @@ -85,16 +112,17 @@ def ci_state(): class TestCISkippedChecksStateField: - def test_ci_skipped_checks_in_ci_integration_state(self): """ci_skipped_checks must be a field in CIIntegrationState.""" from forge.workflow.base import CIIntegrationState + assert "ci_skipped_checks" in CIIntegrationState.__annotations__ def test_initial_feature_state_has_empty_skipped_checks(self): """Fresh feature state initialises ci_skipped_checks to [].""" from forge.models.workflow import TicketType from forge.workflow.feature.state import create_initial_feature_state + state = create_initial_feature_state( thread_id="t", ticket_key="TEST-1", ticket_type=TicketType.FEATURE ) @@ -104,6 +132,7 @@ def test_initial_bug_state_has_empty_skipped_checks(self): """Fresh bug state initialises ci_skipped_checks to [].""" from forge.models.workflow import TicketType from forge.workflow.bug.state import create_initial_bug_state + state = create_initial_bug_state( thread_id="t", ticket_key="TEST-2", ticket_type=TicketType.BUG ) @@ -114,11 +143,8 @@ def test_initial_bug_state_has_empty_skipped_checks(self): class TestWorkerSkipGateDetection: - @pytest.mark.asyncio - async def test_skip_gate_adds_check_to_skipped_list( - self, worker, base_message, ci_state - ): + async def test_skip_gate_adds_check_to_skipped_list(self, worker, base_message, ci_state): """/forge skip-gate appends the check name to ci_skipped_checks.""" msg = _skip_gate_message(base_message, "epoxy") @@ -128,9 +154,7 @@ async def test_skip_gate_adds_check_to_skipped_list( assert "epoxy" in result.get("ci_skipped_checks", []) @pytest.mark.asyncio - async def test_skip_gate_routes_to_ci_evaluator( - self, worker, base_message, ci_state - ): + async def test_skip_gate_routes_to_ci_evaluator(self, worker, base_message, ci_state): """/forge skip-gate unpauses and routes to ci_evaluator.""" msg = _skip_gate_message(base_message, "epoxy") @@ -156,9 +180,7 @@ async def test_unskip_gate_removes_check_from_skipped_list( assert "flamingo" in skipped @pytest.mark.asyncio - async def test_skip_gate_deduplicates( - self, worker, base_message, ci_state - ): + async def test_skip_gate_deduplicates(self, worker, base_message, ci_state): """Skipping the same check twice doesn't add a duplicate.""" ci_state["ci_skipped_checks"] = ["epoxy"] msg = _skip_gate_message(base_message, "epoxy") @@ -169,9 +191,7 @@ async def test_skip_gate_deduplicates( assert result["ci_skipped_checks"].count("epoxy") == 1 @pytest.mark.asyncio - async def test_skip_gate_ignored_outside_ci_stages( - self, worker, base_message - ): + async def test_skip_gate_ignored_outside_ci_stages(self, worker, base_message): """/forge skip-gate has no effect when workflow is not at a CI stage.""" planning_state = make_workflow_state( current_node="prd_approval_gate", @@ -185,9 +205,7 @@ async def test_skip_gate_ignored_outside_ci_stages( assert result.get("is_paused") is True # unchanged @pytest.mark.asyncio - async def test_skip_gate_posts_feedback( - self, worker, base_message, ci_state - ): + async def test_skip_gate_posts_feedback(self, worker, base_message, ci_state): """/forge skip-gate calls _post_skip_gate_feedback.""" msg = _skip_gate_message(base_message, "epoxy") mock_feedback = AsyncMock() @@ -198,22 +216,9 @@ async def test_skip_gate_posts_feedback( mock_feedback.assert_called_once() @pytest.mark.asyncio - async def test_case_insensitive_command_detection( - self, worker, base_message, ci_state - ): + async def test_case_insensitive_command_detection(self, worker, base_message, ci_state): """Command prefix matching is case-insensitive.""" - msg = _skip_gate_message(base_message, "epoxy") - msg = QueueMessage( - message_id=msg.message_id, - event_id=msg.event_id, - source=msg.source, - event_type=msg.event_type, - ticket_key=msg.ticket_key, - payload={ - **msg.payload, - "comment": {"body": "/FORGE SKIP-GATE epoxy"}, - }, - ) + msg = _comment_message(base_message, "/FORGE SKIP-GATE epoxy") with patch.object(worker, "_post_skip_gate_feedback", AsyncMock()): result = await worker._handle_resume_event(msg, ci_state) @@ -225,33 +230,39 @@ async def test_case_insensitive_command_detection( class TestPostSkipGateFeedback: - @pytest.mark.asyncio async def test_posts_github_reply_and_jira_comment(self): """Posts a GitHub PR comment and a Jira audit comment.""" worker = OrchestratorWorker(consumer_name="test") - mock_github = MagicMock() - mock_github.create_issue_comment = AsyncMock() - mock_github.close = AsyncMock() + repo_ref = RepositoryRef( + id="org/repo", + provider=Provider.GITHUB, + connection="default-github", + namespace="org/repo", + default_branch="main", + change_request_mode="fork", + ) + mock_adapter = AsyncMock() mock_jira = MagicMock() mock_jira.add_comment = AsyncMock() mock_jira.close = AsyncMock() - with patch("forge.orchestrator.worker.GitHubClient", return_value=mock_github), \ - patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira): + with ( + patch("forge.orchestrator.worker.get_adapter", return_value=(repo_ref, mock_adapter)), + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + ): await worker._post_skip_gate_feedback( ticket_key="TEST-123", - owner="org", - repo="repo", + repo_ref=repo_ref, pr_number=42, check_name="epoxy", sender="eshulman2", action="skip", ) - mock_github.create_issue_comment.assert_called_once() + mock_adapter.create_comment.assert_called_once() mock_jira.add_comment.assert_called_once() @pytest.mark.asyncio @@ -259,35 +270,89 @@ async def test_unskip_posts_different_message(self): """Unskip action produces a different confirmation message.""" worker = OrchestratorWorker(consumer_name="test") - mock_github = MagicMock() - mock_github.create_issue_comment = AsyncMock() - mock_github.close = AsyncMock() + repo_ref = RepositoryRef( + id="org/repo", + provider=Provider.GITHUB, + connection="default-github", + namespace="org/repo", + default_branch="main", + change_request_mode="fork", + ) + mock_adapter = AsyncMock() mock_jira = MagicMock() mock_jira.add_comment = AsyncMock() mock_jira.close = AsyncMock() - with patch("forge.orchestrator.worker.GitHubClient", return_value=mock_github), \ - patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira): + with ( + patch("forge.orchestrator.worker.get_adapter", return_value=(repo_ref, mock_adapter)), + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + ): await worker._post_skip_gate_feedback( ticket_key="TEST-123", - owner="org", - repo="repo", + repo_ref=repo_ref, pr_number=42, check_name="epoxy", sender="eshulman2", action="unskip", ) - comment = mock_github.create_issue_comment.call_args[0][3] + comment = mock_adapter.create_comment.call_args[0][2] assert "unskip" in comment.lower() or "removed" in comment.lower() # ── CI evaluator: filtering skipped checks ──────────────────────────────────── -class TestEvaluateCIStatusSkipsChecks: +_STATUS_MAP = { + "completed": CheckStatus.COMPLETED, + "in_progress": CheckStatus.IN_PROGRESS, + "pending": CheckStatus.QUEUED, +} +_CONCLUSION_MAP = { + "success": CheckConclusion.SUCCESS, + "failure": CheckConclusion.FAILURE, + "skipped": CheckConclusion.SKIPPED, + "neutral": CheckConclusion.NEUTRAL, + None: CheckConclusion.NONE, +} + + +def _check_run(name: str, status: str, conclusion: str | None) -> CheckRun: + return CheckRun( + name=name, + status=_STATUS_MAP[status], + conclusion=_CONCLUSION_MAP[conclusion], + ) + + +def _mock_adapter_with_checks(checks: list[CheckRun], repo="org/repo", pr_number=42): + """Patch target for ci_evaluator.get_adapter returning a fixed check list.""" + adapter = AsyncMock() + adapter.get_change_request = AsyncMock( + return_value=ChangeRequest( + identity=ChangeRequestIdentity(connection="c", repository_id=repo, native_id=pr_number), + url=f"https://github.com/{repo}/pull/{pr_number}", + title="t", + body="", + state=ChangeRequestState.OPEN, + source_branch="feature", + target_branch="main", + ) + ) + adapter.get_checks = AsyncMock(return_value=checks) + repo_ref = RepositoryRef( + id=repo, + provider=Provider.GITHUB, + connection="c", + namespace=repo, + default_branch="main", + change_request_mode="fork", + ) + return repo_ref, adapter + +class TestEvaluateCIStatusSkipsChecks: @pytest.mark.asyncio async def test_skipped_check_does_not_count_as_failure(self): """A check whose name matches a ci_skipped_checks entry is treated as passing.""" @@ -295,21 +360,24 @@ async def test_skipped_check_does_not_count_as_failure(self): state = make_workflow_state( current_node="ci_evaluator", + current_repo="org/repo", + current_pr_number=42, pr_urls=["https://github.com/org/repo/pull/42"], ci_skipped_checks=["epoxy"], ) - mock_github = MagicMock() - mock_github.get_pull_request = AsyncMock(return_value={"head": {"sha": "abc"}}) - mock_github.get_check_runs = AsyncMock(return_value=[ - {"name": "Run acceptance tests against OpenStack epoxy", - "status": "completed", "conclusion": "failure"}, - {"name": "Run acceptance tests against OpenStack flamingo", - "status": "completed", "conclusion": "success"}, - ]) - mock_github.close = AsyncMock() - - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + repo_ref, adapter = _mock_adapter_with_checks( + [ + _check_run("Run acceptance tests against OpenStack epoxy", "completed", "failure"), + _check_run( + "Run acceptance tests against OpenStack flamingo", "completed", "success" + ), + ] + ) + + with patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(repo_ref, adapter) + ): result = await evaluate_ci_status(state) # Epoxy is skipped, flamingo passed — CI should be "passed" @@ -322,21 +390,24 @@ async def test_all_skipped_checks_plus_pass_routes_to_human_review(self): state = make_workflow_state( current_node="ci_evaluator", + current_repo="org/repo", + current_pr_number=42, pr_urls=["https://github.com/org/repo/pull/42"], ci_skipped_checks=["epoxy", "flamingo"], ) - mock_github = MagicMock() - mock_github.get_pull_request = AsyncMock(return_value={"head": {"sha": "abc"}}) - mock_github.get_check_runs = AsyncMock(return_value=[ - {"name": "Run acceptance tests against OpenStack epoxy", - "status": "completed", "conclusion": "failure"}, - {"name": "Run acceptance tests against OpenStack flamingo", - "status": "completed", "conclusion": "failure"}, - ]) - mock_github.close = AsyncMock() - - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + repo_ref, adapter = _mock_adapter_with_checks( + [ + _check_run("Run acceptance tests against OpenStack epoxy", "completed", "failure"), + _check_run( + "Run acceptance tests against OpenStack flamingo", "completed", "failure" + ), + ] + ) + + with patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(repo_ref, adapter) + ): result = await evaluate_ci_status(state) assert result["ci_status"] == "passed" @@ -349,21 +420,22 @@ async def test_skipped_check_not_in_failed_checks(self): state = make_workflow_state( current_node="ci_evaluator", + current_repo="org/repo", + current_pr_number=42, pr_urls=["https://github.com/org/repo/pull/42"], ci_skipped_checks=["epoxy"], ) - mock_github = MagicMock() - mock_github.get_pull_request = AsyncMock(return_value={"head": {"sha": "abc"}}) - mock_github.get_check_runs = AsyncMock(return_value=[ - {"name": "Run acceptance tests against OpenStack epoxy", - "status": "completed", "conclusion": "failure"}, - {"name": "unit-tests", - "status": "completed", "conclusion": "failure"}, - ]) - mock_github.close = AsyncMock() - - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + repo_ref, adapter = _mock_adapter_with_checks( + [ + _check_run("Run acceptance tests against OpenStack epoxy", "completed", "failure"), + _check_run("unit-tests", "completed", "failure"), + ] + ) + + with patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(repo_ref, adapter) + ): result = await evaluate_ci_status(state) failed = [c["name"] for c in result.get("ci_failed_checks", [])] @@ -377,19 +449,21 @@ async def test_substring_match_is_case_insensitive(self): state = make_workflow_state( current_node="ci_evaluator", + current_repo="org/repo", + current_pr_number=42, pr_urls=["https://github.com/org/repo/pull/42"], ci_skipped_checks=["EPOXY"], # uppercase skip ) - mock_github = MagicMock() - mock_github.get_pull_request = AsyncMock(return_value={"head": {"sha": "abc"}}) - mock_github.get_check_runs = AsyncMock(return_value=[ - {"name": "Run acceptance tests against OpenStack epoxy", - "status": "completed", "conclusion": "failure"}, - ]) - mock_github.close = AsyncMock() + repo_ref, adapter = _mock_adapter_with_checks( + [ + _check_run("Run acceptance tests against OpenStack epoxy", "completed", "failure"), + ] + ) - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(repo_ref, adapter) + ): result = await evaluate_ci_status(state) assert result["ci_status"] == "passed" @@ -405,24 +479,26 @@ async def test_tide_is_ignored_as_permanent_pending_check(self): state = make_workflow_state( current_node="ci_evaluator", + current_repo="org/repo", + current_pr_number=42, pr_urls=["https://github.com/org/repo/pull/42"], ci_skipped_checks=["e2e-openstack"], ) - mock_github = MagicMock() - mock_github.get_pull_request = AsyncMock(return_value={"head": {"sha": "abc"}}) - mock_github.get_check_runs = AsyncMock(return_value=[ - # Openstack e2e Prow checks — skipped by human override - {"name": "ci/prow/e2e-openstack-ovn", - "status": "completed", "conclusion": "failure"}, - # tide — always pending, explicitly filtered by name - {"name": "tide", "status": "pending", "conclusion": None}, - # Real check that passed - {"name": "ci/prow/unit", "status": "completed", "conclusion": "success"}, - ]) - mock_github.close = AsyncMock() - - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + repo_ref, adapter = _mock_adapter_with_checks( + [ + # Openstack e2e Prow checks — skipped by human override + _check_run("ci/prow/e2e-openstack-ovn", "completed", "failure"), + # tide — always pending, explicitly filtered by name + _check_run("tide", "pending", None), + # Real check that passed + _check_run("ci/prow/unit", "completed", "success"), + ] + ) + + with patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(repo_ref, adapter) + ): result = await evaluate_ci_status(state) # e2e-openstack skipped, tide ignored, unit passed → CI passes @@ -436,21 +512,23 @@ async def test_real_pending_check_still_blocks_evaluation(self): state = make_workflow_state( current_node="ci_evaluator", + current_repo="org/repo", + current_pr_number=42, pr_urls=["https://github.com/org/repo/pull/42"], ci_skipped_checks=["e2e-openstack"], ) - mock_github = MagicMock() - mock_github.get_pull_request = AsyncMock(return_value={"head": {"sha": "abc"}}) - mock_github.get_check_runs = AsyncMock(return_value=[ - {"name": "ci/prow/e2e-openstack-ovn", - "status": "completed", "conclusion": "failure"}, - # golint still running — real check, must block - {"name": "ci/prow/golint", "status": "in_progress", "conclusion": None}, - ]) - mock_github.close = AsyncMock() - - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + repo_ref, adapter = _mock_adapter_with_checks( + [ + _check_run("ci/prow/e2e-openstack-ovn", "completed", "failure"), + # golint still running — real check, must block + _check_run("ci/prow/golint", "in_progress", None), + ] + ) + + with patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(repo_ref, adapter) + ): result = await evaluate_ci_status(state) # golint not done → still pending, don't declare passed yet @@ -463,18 +541,21 @@ async def test_empty_skipped_checks_behaves_normally(self): state = make_workflow_state( current_node="ci_evaluator", + current_repo="org/repo", + current_pr_number=42, pr_urls=["https://github.com/org/repo/pull/42"], ci_skipped_checks=[], ) - mock_github = MagicMock() - mock_github.get_pull_request = AsyncMock(return_value={"head": {"sha": "abc"}}) - mock_github.get_check_runs = AsyncMock(return_value=[ - {"name": "unit-tests", "status": "completed", "conclusion": "failure"}, - ]) - mock_github.close = AsyncMock() + repo_ref, adapter = _mock_adapter_with_checks( + [ + _check_run("unit-tests", "completed", "failure"), + ] + ) - with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): + with patch( + "forge.workflow.nodes.ci_evaluator.get_adapter", return_value=(repo_ref, adapter) + ): result = await evaluate_ci_status(state) assert result["ci_status"] == "fixing" diff --git a/tests/unit/workflow/test_concurrent_gate.py b/tests/unit/workflow/test_concurrent_gate.py index d16125a50..3047520ff 100644 --- a/tests/unit/workflow/test_concurrent_gate.py +++ b/tests/unit/workflow/test_concurrent_gate.py @@ -1,6 +1,12 @@ -"""Smoke tests for the concurrent CI/review gate design.""" +"""Smoke tests for the concurrent CI/review gate's routing-table shape. + +These only exercise routing functions with static state and graph compilation; +they do not interleave a CI event and a review event against shared state. For +a regression test covering the actual concurrent-webhook scenario, see +test_review_arriving_during_in_flight_ci_cycle_is_not_dropped in +tests/unit/orchestrator/test_worker.py. +""" -import pytest from langgraph.graph import END diff --git a/tests/unit/workflow/test_implement_review.py b/tests/unit/workflow/test_implement_review.py index aeadc9868..50501f1e1 100644 --- a/tests/unit/workflow/test_implement_review.py +++ b/tests/unit/workflow/test_implement_review.py @@ -7,12 +7,25 @@ import pytest from langgraph.graph import END +from forge.integrations.source_control.contracts import Provider, RepositoryRef from forge.models.workflow import TicketType from forge.workflow.bug.graph import build_bug_graph from forge.workflow.feature.graph import build_feature_graph from forge.workflow.task_takeover.graph import build_task_takeover_graph from tests.fixtures.workflow_states import make_workflow_state + +def _repo_ref(repo: str = "org/repo") -> RepositoryRef: + return RepositoryRef( + id=repo, + provider=Provider.GITHUB, + connection="c", + namespace=repo, + default_branch="main", + change_request_mode="fork", + ) + + # ── State fields ────────────────────────────────────────────────────────────── @@ -367,9 +380,7 @@ async def test_posts_addressing_review_comment_when_review_work_starts(self, tmp mock_git = MagicMock() mock_git._run_git.return_value = MagicMock(stdout="") - mock_github = MagicMock() - mock_github.create_issue_comment = AsyncMock() - mock_github.close = AsyncMock() + mock_adapter = AsyncMock() mock_runner = MagicMock() mock_runner.run = AsyncMock() @@ -392,20 +403,20 @@ async def test_posts_addressing_review_comment_when_review_work_starts(self, tmp "forge.workflow.nodes.implement_review._fetch_pr_review_comments", new=AsyncMock(return_value="# PR Review Feedback\n"), ), - patch("forge.workflow.nodes.implement_review.GitHubClient", return_value=mock_github), + patch( + "forge.workflow.nodes.implement_review.get_adapter", + return_value=(_repo_ref(), mock_adapter), + ), patch( "forge.workflow.nodes.implement_review.ContainerRunner", return_value=mock_runner ), ): result = await implement_review(state) - mock_github.create_issue_comment.assert_called_once_with( - "org", - "repo", - 17, - _REVIEW_ADDRESSING_COMMENT, - ) - mock_github.close.assert_called_once() + mock_adapter.create_comment.assert_awaited_once() + call_args = mock_adapter.create_comment.call_args[0] + assert call_args[1].native_id == 17 + assert call_args[2] == _REVIEW_ADDRESSING_COMMENT mock_runner.run.assert_called_once() assert result["current_node"] == "human_review_gate" @@ -414,261 +425,51 @@ async def test_skips_addressing_review_comment_without_pr_number(self): """No PR comment is posted if the workflow has no PR number.""" from forge.workflow.nodes.implement_review import _post_review_addressing_comment - mock_github = MagicMock() - mock_github.create_issue_comment = AsyncMock() - - with patch("forge.workflow.nodes.implement_review.GitHubClient", return_value=mock_github): + with patch("forge.workflow.nodes.implement_review.get_adapter") as mock_get_adapter: await _post_review_addressing_comment( ticket_key="TEST-789", - owner="org", - repo="repo", + current_repo="org/repo", pr_number=None, ) - mock_github.create_issue_comment.assert_not_called() + mock_get_adapter.assert_not_called() class TestThreadAwareReviewHandling: @pytest.mark.asyncio async def test_processed_threads_are_excluded_from_next_analysis(self): + from forge.integrations.source_control.contracts import Review, ReviewComment, ReviewState from forge.workflow.nodes.implement_review import _fetch_pr_review_comments - github = MagicMock() - github.get_pull_request_review_threads = AsyncMock( - return_value=[ - { - "thread_id": "already-accepted", - "path": "a.py", - "line": 1, - "comments": [{"comment_id": 10, "body": "Done", "author": "r"}], - }, - { - "thread_id": "new-thread", - "path": "b.py", - "line": 2, - "comments": [{"comment_id": 20, "body": "New", "author": "r"}], - }, - ] - ) - github.close = AsyncMock() - - with patch("forge.workflow.nodes.implement_review.GitHubClient", return_value=github): - result = await _fetch_pr_review_comments( - "org", - "repo", - 7, - "", - review_comments=[{"thread_id": "already-accepted", "disposition": "accept"}], - ) - - assert "already-accepted" not in result - assert "new-thread" in result - - @pytest.mark.asyncio - async def test_contested_thread_settled_while_bots_reply_is_last_comment(self): - """A contested thread stays excluded as long as nothing new happened since Forge replied.""" - from forge.workflow.nodes.implement_review import _fetch_pr_review_comments - - github = MagicMock() - github.get_pull_request_review_threads = AsyncMock( - return_value=[ - { - "thread_id": "contested-thread", - "path": "a.py", - "line": 1, - "comments": [ - { - "comment_id": 10, - "body": "This is wrong", - "author": "reviewer", - "created_at": "2026-01-01T00:00:00Z", - }, - { - "comment_id": 11, - "body": "This conflicts with the public API.", - "author": "forge-bot", - "created_at": "2026-01-02T00:00:00Z", - }, - ], - }, - ] - ) - github.close = AsyncMock() - login_client = MagicMock() - login_client.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - login_client.close = AsyncMock() - - with patch( - "forge.workflow.nodes.implement_review.GitHubClient", - side_effect=[github, login_client], - ): - result = await _fetch_pr_review_comments( - "org", - "repo", - 7, - "", - review_comments=[{"thread_id": "contested-thread", "disposition": "contest"}], - ) - - assert "contested-thread" not in result - - @pytest.mark.asyncio - async def test_contested_thread_resurfaces_once_human_replies_after_forge(self): - """A human reply after Forge's objection re-opens the thread for re-analysis.""" - from forge.workflow.nodes.implement_review import _fetch_pr_review_comments - - github = MagicMock() - github.get_pull_request_review_threads = AsyncMock( + adapter = AsyncMock() + adapter.get_review_thread_comments = AsyncMock( return_value=[ - { - "thread_id": "contested-thread", - "path": "a.py", - "line": 1, - "comments": [ - { - "comment_id": 10, - "body": "This is wrong", - "author": "reviewer", - "created_at": "2026-01-01T00:00:00Z", - }, - { - "comment_id": 11, - "body": "This conflicts with the public API.", - "author": "forge-bot", - "created_at": "2026-01-02T00:00:00Z", - }, - { - "comment_id": 12, - "body": "Please do it anyway.", - "author": "reviewer", - "created_at": "2026-01-03T00:00:00Z", - }, - ], - }, + Review( + id="already-accepted", + state=ReviewState.COMMENTED, + body="", + author="r", + comments=[ReviewComment(id="10", body="Done", author="r", path="a.py", line=1)], + ), + Review( + id="new-thread", + state=ReviewState.COMMENTED, + body="", + author="r", + comments=[ReviewComment(id="20", body="New", author="r", path="b.py", line=2)], + ), ] ) - github.close = AsyncMock() - login_client = MagicMock() - login_client.get_authenticated_user = AsyncMock(return_value={"login": "forge-bot"}) - login_client.close = AsyncMock() with patch( - "forge.workflow.nodes.implement_review.GitHubClient", - side_effect=[github, login_client], + "forge.workflow.nodes.implement_review.get_adapter", + return_value=(_repo_ref(), adapter), ): - result = await _fetch_pr_review_comments( - "org", - "repo", - 7, - "", - review_comments=[{"thread_id": "contested-thread", "disposition": "contest"}], - ) - - assert "contested-thread" in result - assert "Please do it anyway." in result - - @pytest.mark.asyncio - async def test_reply_to_review_threads_uses_skip_addressed_guard(self): - """A second guard: never re-post to a decision already marked addressed.""" - from forge.workflow.nodes.implement_review import _reply_to_review_threads - - with patch( - "forge.workflow.nodes.implement_review.reply_to_review_decisions", - new=AsyncMock(), - ) as reply_mock: - await _reply_to_review_threads( - owner="org", repo="repo", pr_number=9, decisions=[{"thread_id": "t"}] - ) - - reply_mock.assert_awaited_once_with( - repo_full_name="org/repo", - pr_number=9, - decisions=[{"thread_id": "t"}], - skip_addressed=True, - ) - - @pytest.mark.asyncio - async def test_no_bot_login_lookup_when_no_contested_decisions(self): - """Avoid the extra GET /user round trip when there's nothing to settle-check.""" - from forge.workflow.nodes.implement_review import _fetch_pr_review_comments - - github = MagicMock() - github.get_pull_request_review_threads = AsyncMock( - return_value=[ - { - "thread_id": "new-thread", - "path": "b.py", - "line": 2, - "comments": [ - { - "comment_id": 20, - "body": "New", - "author": "r", - "created_at": "2026-01-01T00:00:00Z", - } - ], - }, - ] - ) - github.close = AsyncMock() - - with patch( - "forge.workflow.nodes.implement_review.GitHubClient", return_value=github - ) as ctor: - result = await _fetch_pr_review_comments("org", "repo", 7, "", review_comments=[]) + result = await _fetch_pr_review_comments("org/repo", 7, "", {"already-accepted"}) - ctor.assert_called_once() + assert "already-accepted" not in result assert "new-thread" in result - @pytest.mark.asyncio - async def test_contested_thread_resurfaces_when_bot_login_lookup_fails(self): - """A failed identity lookup must fail toward re-analysis, not a coincidental - empty-string match against a deleted account's comment author.""" - from forge.workflow.nodes.implement_review import _fetch_pr_review_comments - - github = MagicMock() - github.get_pull_request_review_threads = AsyncMock( - return_value=[ - { - "thread_id": "contested-thread", - "path": "a.py", - "line": 1, - "comments": [ - { - "comment_id": 10, - "body": "This is wrong", - "author": "reviewer", - "created_at": "2026-01-01T00:00:00Z", - }, - { - "comment_id": 11, - "body": "This conflicts with the public API.", - "author": "", # e.g. a deleted account, per GitHub's API - "created_at": "2026-01-02T00:00:00Z", - }, - ], - }, - ] - ) - github.close = AsyncMock() - login_client = MagicMock() - login_client.get_authenticated_user = AsyncMock(side_effect=RuntimeError("auth failed")) - login_client.close = AsyncMock() - - with patch( - "forge.workflow.nodes.implement_review.GitHubClient", - side_effect=[github, login_client], - ): - result = await _fetch_pr_review_comments( - "org", - "repo", - 7, - "", - review_comments=[{"thread_id": "contested-thread", "disposition": "contest"}], - ) - - assert "contested-thread" in result - @pytest.mark.asyncio async def test_legacy_objections_file_still_pauses_for_response(self, tmp_path): from forge.workflow.nodes.implement_review import implement_review @@ -790,81 +591,14 @@ async def run_container(**_kwargs): assert mock_runner.run.await_count == 2 assert result["current_node"] == "review_response_gate" - # Once Forge has replied to the contested thread, it's marked addressed - # so the skip_addressed guard in reply_to_review_decisions can catch a - # stray re-reply if the same decision ever resurfaces unchanged. - assert result["contested_comments"] == [{**decisions[1], "status": "addressed"}] + assert result["contested_comments"] == [decisions[1]] assert result["review_comments"][0] == { **decisions[0], "response": "Forge verified this feedback; no additional code change was needed.", } - assert result["review_comments"][1] == {**decisions[1], "status": "addressed"} + assert result["review_comments"][1] == decisions[1] assert reply_threads.await_count == 2 - @pytest.mark.asyncio - async def test_contested_decision_marked_addressed_after_reply(self, tmp_path): - """Once Forge has replied to a contested thread, the persisted decision - carries status=addressed, so a decision that somehow re-surfaces - unchanged is skipped by reply_to_review_decisions' skip_addressed guard.""" - from forge.workflow.nodes.implement_review import implement_review - - decisions = [ - { - "thread_id": "contested-thread", - "comment_id": 20, - "disposition": "contest", - "feedback": "", - "reason": "Conflicts with the public API", - "response": "This conflicts with the documented public API. Can you confirm?", - }, - ] - - async def run_container(**_kwargs): - if not (tmp_path / ".forge" / "review-decisions.json").exists(): - (tmp_path / ".forge" / "review-decisions.json").write_text(json.dumps(decisions)) - (tmp_path / ".forge" / "review-plan.md").write_text("# No actionable items") - - mock_runner = MagicMock() - mock_runner.run = AsyncMock(side_effect=run_container) - mock_git = MagicMock() - mock_git.has_uncommitted_changes.return_value = False - mock_git._run_git.return_value = MagicMock(stdout="") - state = make_workflow_state( - ticket_key="TEST-234", - current_node="implement_review", - workspace_path=str(tmp_path), - current_repo="org/repo", - current_pr_number=9, - feedback_comment="Review", - context={"branch_name": "forge/TEST-234"}, - ) - - with ( - patch( - "forge.workflow.nodes.implement_review.prepare_workspace", - return_value=(str(tmp_path), mock_git), - ), - patch( - "forge.workflow.nodes.implement_review._fetch_pr_review_comments", - new=AsyncMock(return_value="# Review"), - ), - patch( - "forge.workflow.nodes.implement_review._post_review_addressing_comment", - new=AsyncMock(), - ), - patch( - "forge.workflow.nodes.implement_review._reply_to_review_threads", - new=AsyncMock(), - ), - patch( - "forge.workflow.nodes.implement_review.ContainerRunner", - return_value=mock_runner, - ), - ): - result = await implement_review(state) - - assert result["review_comments"][0]["status"] == "addressed" - def test_confirming_one_thread_routes_to_implementation_with_others_pending(self): from forge.workflow.nodes.implement_review import route_review_response diff --git a/tests/unit/workflow/test_no_concrete_client_imports.py b/tests/unit/workflow/test_no_concrete_client_imports.py new file mode 100644 index 000000000..774679858 --- /dev/null +++ b/tests/unit/workflow/test_no_concrete_client_imports.py @@ -0,0 +1,36 @@ +import pathlib +import re + +SRC_DIR = pathlib.Path(__file__).resolve().parents[3] / "src" / "forge" +WORKFLOW_DIR = SRC_DIR / "workflow" +WORKSPACE_DIR = SRC_DIR / "workspace" + +# Matches any dotted reference to the concrete client module or its parent +# package, e.g. `from forge.integrations.github.client import GitHubClient` +# or `forge.integrations.github.client.GitHubClient(...)`. This is a textual +# scan, not an import graph: it does NOT catch a two-step alias like +# `from forge.integrations import github` followed later by +# `github.client.GitHubClient(...)`, since that indirection never spells out +# `integrations.github` as one dotted token. It also only scans +# WORKFLOW_DIR/WORKSPACE_DIR below, not other modules (e.g. worker.py) that +# were migrated onto get_adapter. +_CONCRETE_CLIENT_PATTERN = re.compile(r"integrations\.github\.client|integrations\.github\b") + + +def _find_offenders(directory: pathlib.Path) -> list[str]: + offenders = [] + for path in directory.rglob("*.py"): + text = path.read_text() + if _CONCRETE_CLIENT_PATTERN.search(text): + offenders.append(str(path.relative_to(directory))) + return offenders + + +def test_no_workflow_module_imports_github_client(): + offenders = _find_offenders(WORKFLOW_DIR) + assert not offenders, f"workflow modules still import the concrete GitHub client: {offenders}" + + +def test_no_workspace_module_imports_github_client(): + offenders = _find_offenders(WORKSPACE_DIR) + assert not offenders, f"workspace modules still import the concrete GitHub client: {offenders}" diff --git a/tests/unit/workflow/test_pr_state.py b/tests/unit/workflow/test_pr_state.py index e3cf8892f..0aab124ef 100644 --- a/tests/unit/workflow/test_pr_state.py +++ b/tests/unit/workflow/test_pr_state.py @@ -1,3 +1,15 @@ +from datetime import UTC, datetime + +from forge.integrations.source_control.contracts import ( + Actor, + ChangeRequest, + ChangeRequestIdentity, + ChangeRequestState, + EventKind, + NormalizedEvent, + Provider, + RepositoryRef, +) from forge.workflow.pr_state import ( activate_pull_request_for_event, all_pull_requests_merged, @@ -6,6 +18,40 @@ ) +def _event( + repo="acme/payments", + native_id=42, + url="https://github.com/acme/payments/pull/42", +) -> NormalizedEvent: + repo_ref = RepositoryRef( + id=repo, + provider=Provider.GITHUB, + connection="default-github", + namespace=repo, + default_branch="main", + change_request_mode="fork", + ) + return NormalizedEvent( + id="e1", + kind=EventKind.CR_UPDATED, + repo_ref=repo_ref, + actor=Actor(login="octocat", is_bot=False), + received_at=datetime(2026, 1, 1, tzinfo=UTC), + change_request=ChangeRequest( + identity=ChangeRequestIdentity( + connection="default-github", repository_id=repo, native_id=native_id + ), + url=url, + title="t", + body="", + state=ChangeRequestState.OPEN, + source_branch="feature", + target_branch="main", + draft=False, + ), + ) + + def _multi_repo_state() -> dict: return { "current_repo": "acme/frontend", @@ -16,7 +62,7 @@ def _multi_repo_state() -> dict: "ci_status": "passed", "pr_merged": False, "pull_requests": { - "acme/backend": { + "acme/backend:10": { "repo": "acme/backend", "number": 10, "url": "https://github.com/acme/backend/pull/10", @@ -28,7 +74,7 @@ def _multi_repo_state() -> dict: "merged": False, "lifecycle_node": "review_response_gate", }, - "acme/frontend": { + "acme/frontend:20": { "repo": "acme/frontend", "number": 20, "url": "https://github.com/acme/frontend/pull/20", @@ -42,111 +88,163 @@ def _multi_repo_state() -> dict: } -def test_github_event_activates_matching_repo_pr() -> None: - state = _multi_repo_state() - payload = { - "repository": {"full_name": "acme/backend"}, - "pull_request": {"number": 10}, - } +# ── Brief Step 1: core NormalizedEvent keying ────────────────────────────── - 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_activate_pull_request_for_event_keys_by_repo_and_number() -> None: + event = _event() + key = "acme/payments:42" + state = {"pull_requests": {key: {"number": 42, "url": event.change_request.url}}} + activated = activate_pull_request_for_event(state, event) -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 activated["current_pr_number"] == 42 + assert activated["current_repo"] == "acme/payments" + + +def test_activate_pull_request_for_event_returns_state_unchanged_for_none_event() -> None: + state = {"pull_requests": {}} + + assert activate_pull_request_for_event(state, None) is state + + +def test_event_targets_pull_request_true_for_matching_key() -> None: + event = _event() + key = "acme/payments:42" + state = {"pull_requests": {key: {"number": 42, "url": event.change_request.url}}} + + assert event_targets_pull_request(state, event) is True + + +def test_event_targets_pull_request_false_for_none_event() -> None: + assert event_targets_pull_request({"pull_requests": {}}, None) is False + + +def test_two_prs_same_repo_get_distinct_keys() -> None: + """The whole point of keying by repo+number instead of owner/repo: multiple + concurrent PRs against the same repo must not collide on one dict slot.""" + state = { + "pull_requests": { + "acme/payments:42": { + "number": 42, + "url": "https://github.com/acme/payments/pull/42", + "ci_status": "passed", + }, + "acme/payments:43": { + "number": 43, + "url": "https://github.com/acme/payments/pull/43", + "ci_status": "fixing", + }, + } } + event_42 = _event(native_id=42, url="https://github.com/acme/payments/pull/42") + event_43 = _event(native_id=43, url="https://github.com/acme/payments/pull/43") - assert activate_pull_request_for_event(state, payload) == state + activated_42 = activate_pull_request_for_event(state, event_42) + activated_43 = activate_pull_request_for_event(state, event_43) + assert activated_42["current_pr_number"] == 42 + assert activated_42["ci_status"] == "passed" + assert activated_43["current_pr_number"] == 43 + assert activated_43["ci_status"] == "fixing" -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 +# ── Preserved behavioral coverage (ported from the raw-payload version) ───── -def test_check_run_event_activates_matching_repo_pr() -> None: +def test_event_activates_matching_repo_pr() -> None: state = _multi_repo_state() - payload = { - "repository": {"full_name": "acme/backend"}, - "check_run": {"pull_requests": [{"number": 10}]}, - } + event = _event(repo="acme/backend", native_id=10) - activated = activate_pull_request_for_event(state, payload) + activated = activate_pull_request_for_event(state, event) 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_same_repo_event_preserves_existing_workspace() -> None: +def test_event_for_unknown_pr_does_not_change_active_repo() -> 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) + event = _event(repo="acme/other", native_id=99) - assert activated["workspace_path"] == "/tmp/forge-AISOS-1-active" + assert activate_pull_request_for_event(state, event) == state -def test_issue_comment_event_activates_matching_repo_pr() -> None: +def test_same_repo_event_preserves_existing_workspace() -> None: state = _multi_repo_state() - payload = { - "repository": {"full_name": "acme/backend"}, - "issue": {"number": 10}, - } + state["workspace_path"] = "/tmp/forge-AISOS-1-active" + event = _event(repo="acme/frontend", native_id=20) - activated = activate_pull_request_for_event(state, payload) + activated = activate_pull_request_for_event(state, event) - assert activated["current_repo"] == "acme/backend" - assert activated["current_pr_number"] == 10 + assert activated["workspace_path"] == "/tmp/forge-AISOS-1-active" 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}}, + _event(repo="acme/backend", native_id=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" + assert saved["pull_requests"]["acme/backend:10"]["ci_status"] == "passed" + assert saved["pull_requests"]["acme/backend:10"]["ci_fix_attempt"] == 0 + # The frontend record is left untouched. + assert saved["pull_requests"]["acme/frontend:20"]["ci_status"] == "passed" + + +def test_save_creates_distinct_slot_per_pr_number() -> None: + """Two PRs on the same repo saved from the scalar view land in separate slots.""" + base = _multi_repo_state() + base["pull_requests"] = {} + + first = save_active_pull_request( + { + **base, + "current_repo": "acme/api", + "current_pr_number": 1, + "current_pr_url": "https://github.com/acme/api/pull/1", + "ci_status": "passed", + } + ) + both = save_active_pull_request( + { + **first, + "current_repo": "acme/api", + "current_pr_number": 2, + "current_pr_url": "https://github.com/acme/api/pull/2", + "ci_status": "fixing", + } + ) + + assert both["pull_requests"]["acme/api:1"]["ci_status"] == "passed" + assert both["pull_requests"]["acme/api:2"]["ci_status"] == "fixing" def test_merge_completion_requires_every_pr() -> None: state = _multi_repo_state() - state["pull_requests"]["acme/backend"]["merged"] = True + state["pull_requests"]["acme/backend:10"]["merged"] = True assert not all_pull_requests_merged(state) - state["pull_requests"]["acme/frontend"]["merged"] = True + state["pull_requests"]["acme/frontend:20"]["merged"] = True assert all_pull_requests_merged(state) +# ── Number-unknown (URL-keyed) fallback ──────────────────────────────────── + + 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"] = { + state["pull_requests"]["acme/backend:10"]["merged"] = True + state["pull_requests"]["acme/frontend:20"]["merged"] = True + state["pull_requests"]["acme/docs:https://github.com/acme/docs/pull/30"] = { "repo": "acme/docs", "url": "https://github.com/acme/docs/pull/30", "number": None, @@ -156,43 +254,141 @@ def test_url_only_pr_blocks_aggregate_merge_completion() -> None: assert not all_pull_requests_merged(state) -def test_later_webhook_hydrates_url_only_pr_number() -> None: +def test_save_with_unknown_number_keys_by_url() -> None: + saved = save_active_pull_request( + { + "current_repo": "acme/docs", + "current_pr_number": None, + "current_pr_url": "https://github.com/acme/docs/pull/30", + "pull_requests": {}, + } + ) + + assert "acme/docs:https://github.com/acme/docs/pull/30" in saved["pull_requests"] + assert saved["pull_requests"]["acme/docs:https://github.com/acme/docs/pull/30"]["number"] is None + + +def test_save_without_number_or_url_is_noop() -> None: + state = {"current_repo": "acme/docs", "pull_requests": {}} + assert save_active_pull_request(state) is state + + +def test_later_webhook_hydrates_and_rekeys_url_only_pr() -> None: + url = "https://github.com/acme/docs/pull/30" state = _multi_repo_state() - state["pull_requests"]["acme/docs"] = { + state["pull_requests"][f"acme/docs:{url}"] = { "repo": "acme/docs", - "url": "https://github.com/acme/docs/pull/30", + "url": url, "number": None, "merged": False, } - payload = { - "repository": {"full_name": "acme/docs"}, - "pull_request": { - "number": 30, - "html_url": "https://github.com/acme/docs/pull/30", - }, - } + event = _event(repo="acme/docs", native_id=30, url=url) - activated = activate_pull_request_for_event(state, payload) + activated = activate_pull_request_for_event(state, event) assert activated["current_repo"] == "acme/docs" assert activated["current_pr_number"] == 30 - assert activated["pull_requests"]["acme/docs"]["number"] == 30 + # Record re-keyed from its URL slot to the numbered slot; no duplicate left. + assert "acme/docs:30" in activated["pull_requests"] + assert f"acme/docs:{url}" not in activated["pull_requests"] + assert activated["pull_requests"]["acme/docs:30"]["number"] == 30 + + +def test_rekeyed_pr_does_not_duplicate_on_subsequent_save() -> None: + url = "https://github.com/acme/docs/pull/30" + state = _multi_repo_state() + state["pull_requests"][f"acme/docs:{url}"] = { + "repo": "acme/docs", + "url": url, + "number": None, + "merged": False, + } + activated = activate_pull_request_for_event( + state, _event(repo="acme/docs", native_id=30, url=url) + ) + activated["ci_status"] = "passed" + + saved = save_active_pull_request(activated) + + docs_keys = [k for k in saved["pull_requests"] if k.startswith("acme/docs:")] + assert docs_keys == ["acme/docs:30"] + assert saved["pull_requests"]["acme/docs:30"]["ci_status"] == "passed" def test_url_only_pr_rejects_different_pr_in_same_repo() -> None: + url = "https://github.com/acme/docs/pull/30" state = _multi_repo_state() - state["pull_requests"]["acme/docs"] = { + state["pull_requests"][f"acme/docs:{url}"] = { "repo": "acme/docs", - "url": "https://github.com/acme/docs/pull/30", + "url": url, "number": None, "merged": False, } - payload = { - "repository": {"full_name": "acme/docs"}, - "pull_request": { - "number": 31, - "html_url": "https://github.com/acme/docs/pull/31", + # A different PR (#31) on the same repo must not match the url-only #30 record. + event = _event( + repo="acme/docs", + native_id=31, + url="https://github.com/acme/docs/pull/31", + ) + + assert not event_targets_pull_request(state, event) + + +# ── Legacy bare-repo key fallback ─────────────────────────────────────────── +# Workflows checkpointed mid-CI/mid-review before per-PR keying shipped have +# their record stored under `repo` alone rather than `repo:number`/`repo:url`. + + +def _legacy_state() -> dict: + return { + "current_repo": None, + "current_pr_number": None, + "current_pr_url": None, + "pull_requests": { + "acme/legacy": { + "repo": "acme/legacy", + "number": 99, + "url": "https://github.com/acme/legacy/pull/99", + "ci_status": "pending", + "merged": False, + "lifecycle_node": "ci_evaluator", + }, }, } - assert not event_targets_pull_request(state, payload) + +def test_event_targets_pull_request_matches_legacy_bare_repo_key() -> None: + state = _legacy_state() + event = _event( + repo="acme/legacy", native_id=99, url="https://github.com/acme/legacy/pull/99" + ) + + assert event_targets_pull_request(state, event) + + +def test_activate_pull_request_for_event_hydrates_from_legacy_bare_repo_key() -> None: + state = _legacy_state() + event = _event( + repo="acme/legacy", native_id=99, url="https://github.com/acme/legacy/pull/99" + ) + + activated = activate_pull_request_for_event(state, event) + + assert activated["current_repo"] == "acme/legacy" + assert activated["current_pr_number"] == 99 + assert activated["ci_status"] == "pending" + + +def test_save_migrates_legacy_bare_repo_key_to_numbered_key() -> None: + state = _legacy_state() + event = _event( + repo="acme/legacy", native_id=99, url="https://github.com/acme/legacy/pull/99" + ) + activated = activate_pull_request_for_event(state, event) + activated["ci_status"] = "passed" + + saved = save_active_pull_request(activated) + + assert "acme/legacy" not in saved["pull_requests"] + assert saved["pull_requests"]["acme/legacy:99"]["ci_status"] == "passed" + assert saved["pull_requests"]["acme/legacy:99"]["lifecycle_node"] == "ci_evaluator" diff --git a/tests/unit/workflow/test_yolo_mode.py b/tests/unit/workflow/test_yolo_mode.py index b4a261c14..b05cc8d5b 100644 --- a/tests/unit/workflow/test_yolo_mode.py +++ b/tests/unit/workflow/test_yolo_mode.py @@ -86,7 +86,7 @@ def test_yolo_mode_false_for_github_source(self): from forge.models.events import EventSource msg = MagicMock() msg.ticket_key = "TEST-1" - msg.source = EventSource.GITHUB + msg.source = EventSource.SOURCE_CONTROL msg.event_type = "pull_request" msg.event_id = "evt-1" msg.retry_count = 0 diff --git a/tests/unit/workflow/utils/test_review_decisions.py b/tests/unit/workflow/utils/test_review_decisions.py index 4f376c749..d1573fff9 100644 --- a/tests/unit/workflow/utils/test_review_decisions.py +++ b/tests/unit/workflow/utils/test_review_decisions.py @@ -1,7 +1,8 @@ -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, patch import pytest +from forge.integrations.source_control.contracts import Provider, RepositoryRef, ReviewComment from forge.workflow.utils.review_decisions import ( decision_matches_comment, flatten_review_threads, @@ -10,6 +11,17 @@ ) +def _repo_ref(repo: str = "org/repo") -> RepositoryRef: + return RepositoryRef( + id=repo, + provider=Provider.GITHUB, + connection="c", + namespace=repo, + default_branch="main", + change_request_mode="fork", + ) + + def test_merge_keeps_latest_decision_per_thread() -> None: previous = [ {"thread_id": "thread-a", "disposition": "reply"}, @@ -49,16 +61,33 @@ def test_decision_matches_original_or_forge_reply_comment() -> None: assert not decision_matches_comment(decision, 12) +def test_decision_matches_comment_across_int_and_str_ids() -> None: + """forge_reply_id now comes back as a str from the adapter; comment_id may + still be int from older persisted state. Comparison must not care which.""" + decision = {"comment_id": 10, "forge_reply_id": "11"} + + assert decision_matches_comment(decision, "10") + assert decision_matches_comment(decision, 11) + assert not decision_matches_comment(decision, 12) + + @pytest.mark.asyncio async def test_reply_records_forge_comment_id() -> None: - github = MagicMock() - github.reply_to_review_comment = AsyncMock(return_value={"id": 11}) - github.close = AsyncMock() + adapter = AsyncMock() + adapter.reply_to_comment = AsyncMock( + return_value=ReviewComment(id="11", body="ok", author="forge") + ) decision = {"comment_id": 10, "response": "Please confirm."} - with patch("forge.integrations.github.client.GitHubClient", return_value=github): - await reply_to_review_decisions( - repo_full_name="org/repo", pr_number=7, decisions=[decision] - ) - - assert decision["forge_reply_id"] == 11 + with patch( + "forge.workflow.utils.review_decisions.get_adapter", + return_value=(_repo_ref(), adapter), + ): + await reply_to_review_decisions(current_repo="org/repo", pr_number=7, decisions=[decision]) + + assert decision["forge_reply_id"] == "11" + adapter.reply_to_comment.assert_awaited_once() + call_args = adapter.reply_to_comment.call_args[0] + assert call_args[1].native_id == 7 + assert call_args[2] == "10" + assert call_args[3] == "Please confirm." diff --git a/tests/unit/workflow/utils/test_source_control_resolver.py b/tests/unit/workflow/utils/test_source_control_resolver.py new file mode 100644 index 000000000..ce74dcf2b --- /dev/null +++ b/tests/unit/workflow/utils/test_source_control_resolver.py @@ -0,0 +1,46 @@ +import pytest + +from forge.integrations.source_control.contracts import ( + ChangeRequestIdentity, + RepositoryRef, + Provider, +) +from forge.integrations.source_control.errors import NotFoundError +from forge.workflow.utils.source_control import get_adapter, identity_for + + +def test_get_adapter_returns_repo_ref_and_adapter(): + import forge.config as config_module + import forge.integrations.source_control.github # noqa: F401 (register factory) + import forge.integrations.source_control.registry as registry_module + + registry_module.get_registry.cache_clear() + config_module.get_settings.cache_clear() + try: + repo_ref, adapter = get_adapter("acme/widgets") + assert repo_ref.namespace == "acme/widgets" + assert repo_ref.provider == Provider.GITHUB + assert adapter is not None + finally: + registry_module.get_registry.cache_clear() + config_module.get_settings.cache_clear() + + +def test_get_adapter_rejects_empty_identifier(): + with pytest.raises(NotFoundError): + get_adapter("") + + +def test_identity_for_builds_composite_identity(): + repo_ref = RepositoryRef( + id="acme/widgets", + provider=Provider.GITHUB, + connection="github-default", + namespace="acme/widgets", + default_branch="main", + change_request_mode="fork", + ) + identity = identity_for(repo_ref, 42) + assert identity == ChangeRequestIdentity( + connection="github-default", repository_id="acme/widgets", native_id=42 + ) diff --git a/tests/unit/workspace/test_git_ops_commit.py b/tests/unit/workspace/test_git_ops_commit.py index 1843c3230..fac11c947 100644 --- a/tests/unit/workspace/test_git_ops_commit.py +++ b/tests/unit/workspace/test_git_ops_commit.py @@ -4,9 +4,12 @@ from pathlib import Path from unittest.mock import MagicMock, patch +from forge.integrations.source_control.contracts import GitCredentials from forge.workspace.git_ops import GitOperations from forge.workspace.manager import Workspace +_CREDENTIALS = GitCredentials(host="github.com", token="test-token") + def _run_git(repo: Path, *args: str) -> subprocess.CompletedProcess[str]: return subprocess.run( @@ -39,7 +42,7 @@ def test_commit_supplies_identity_without_global_git_config(tmp_path, monkeypatc ticket_key="TEST-123", ) with patch("forge.workspace.git_ops.get_settings", return_value=settings): - git = GitOperations(workspace) + git = GitOperations(workspace, _CREDENTIALS) empty_global_config = tmp_path / "empty-gitconfig" empty_global_config.touch() @@ -81,7 +84,7 @@ def test_commit_ignores_untracked_forge_metadata(tmp_path): ticket_key="TEST-123", ) with patch("forge.workspace.git_ops.get_settings", return_value=settings): - git = GitOperations(workspace) + git = GitOperations(workspace, _CREDENTIALS) assert git.commit("Test commit") is False @@ -100,7 +103,7 @@ def test_has_uncommitted_changes_excludes_forge_metadata(tmp_path): ticket_key="TEST-123", ) with patch("forge.workspace.git_ops.get_settings", return_value=settings): - git = GitOperations(workspace) + git = GitOperations(workspace, _CREDENTIALS) forge_dir = repo / ".forge" forge_dir.mkdir() diff --git a/tests/unit/workspace/test_git_ops_credentials.py b/tests/unit/workspace/test_git_ops_credentials.py new file mode 100644 index 000000000..c32c4c28e --- /dev/null +++ b/tests/unit/workspace/test_git_ops_credentials.py @@ -0,0 +1,113 @@ +"""Tests for GitOperations' use of connection-scoped GitCredentials. + +GitOperations must never hardcode github.com or the process-wide +GITHUB_TOKEN -- every clone/remote URL and TLS trust setting is derived +from the GitCredentials passed in at construction, so operations against a +non-default connection (GitHub Enterprise, or a second org with its own +token) hit the right host with the right credential. +""" + +from pathlib import Path +from unittest.mock import MagicMock, patch + +from forge.integrations.source_control.contracts import GitCredentials +from forge.workspace.git_ops import GitOperations +from forge.workspace.manager import Workspace + + +def _git_ops(tmp_path: Path, credentials: GitCredentials) -> GitOperations: + workspace = Workspace( + path=tmp_path / "repo", + repo_name="acme/widgets", + branch_name="forge/test", + ticket_key="TEST-1", + ) + with patch("forge.workspace.git_ops.get_settings", return_value=MagicMock()): + return GitOperations(workspace, credentials) + + +def test_clone_builds_url_from_enterprise_host_not_github_com(tmp_path): + """A GitHub Enterprise connection's host must be used for the default + clone URL, not a hardcoded github.com.""" + credentials = GitCredentials(host="ghe.example.com", token="ghe-token") + git = _git_ops(tmp_path, credentials) + + with patch("forge.workspace.git_ops.subprocess.run") as run: + run.return_value = MagicMock(returncode=0, stderr="") + git.clone() + + cmd = run.call_args.args[0] + assert "https://x-access-token:ghe-token@ghe.example.com/acme/widgets.git" in cmd + + +def test_add_fork_remote_builds_url_from_credentials_host(tmp_path): + """The fork remote URL must use this workspace's connection host/token, + not the process-wide GITHUB_TOKEN against github.com.""" + credentials = GitCredentials(host="ghe.example.com", token="ghe-token") + git = _git_ops(tmp_path, credentials) + + with patch.object(git, "_run_git") as run_git: + run_git.return_value.stdout = "" + git.add_fork_remote("forge-bot", "widgets") + + add_call = next(c for c in run_git.call_args_list if c.args[:2] == ("remote", "add")) + fork_url = add_call.args[3] + assert fork_url == "https://x-access-token:ghe-token@ghe.example.com/forge-bot/widgets.git" + + +def test_run_git_sets_ssl_cainfo_when_ca_path_configured(tmp_path): + """A connection's CA bundle must be trusted via GIT_SSL_CAINFO, so a + self-signed Enterprise Server cert doesn't fail TLS verification.""" + credentials = GitCredentials( + host="ghe.example.com", token="ghe-token", ca_path="/etc/ssl/certs/ghe-ca.pem" + ) + git = _git_ops(tmp_path, credentials) + git.repo_path.mkdir(parents=True, exist_ok=True) + + with patch("forge.workspace.git_ops.subprocess.run") as run: + run.return_value = MagicMock(returncode=0, stdout="", stderr="") + git._run_git("status") + + assert run.call_args.kwargs["env"]["GIT_SSL_CAINFO"] == "/etc/ssl/certs/ghe-ca.pem" + + +def test_run_git_omits_env_override_when_no_ca_path(tmp_path): + """The common case (no custom CA) must not override the subprocess + environment at all -- inherits the process environment unmodified.""" + credentials = GitCredentials(host="github.com", token="test-token") + git = _git_ops(tmp_path, credentials) + git.repo_path.mkdir(parents=True, exist_ok=True) + + with patch("forge.workspace.git_ops.subprocess.run") as run: + run.return_value = MagicMock(returncode=0, stdout="", stderr="") + git._run_git("status") + + assert run.call_args.kwargs["env"] is None + + +def test_clone_sets_ssl_cainfo_when_ca_path_configured(tmp_path): + """clone() bypasses _run_git (no repo to cwd into yet), so it must + independently apply the same CA trust setting.""" + credentials = GitCredentials( + host="ghe.example.com", token="ghe-token", ca_path="/etc/ssl/certs/ghe-ca.pem" + ) + git = _git_ops(tmp_path, credentials) + + with patch("forge.workspace.git_ops.subprocess.run") as run: + run.return_value = MagicMock(returncode=0, stderr="") + git.clone() + + assert run.call_args.kwargs["env"]["GIT_SSL_CAINFO"] == "/etc/ssl/certs/ghe-ca.pem" + + +def test_explicit_repo_url_is_used_as_is(tmp_path): + """An explicit repo_url overrides the credentials-derived default.""" + credentials = GitCredentials(host="github.com", token="test-token") + git = _git_ops(tmp_path, credentials) + + with patch("forge.workspace.git_ops.subprocess.run") as run: + run.return_value = MagicMock(returncode=0, stderr="") + git.clone(repo_url="https://custom.example.com/some/repo.git") + + cmd = run.call_args.args[0] + assert "https://custom.example.com/some/repo.git" in cmd diff --git a/tests/unit/workspace/test_git_ops_redaction.py b/tests/unit/workspace/test_git_ops_redaction.py index 49d636c4e..fe97b53a5 100644 --- a/tests/unit/workspace/test_git_ops_redaction.py +++ b/tests/unit/workspace/test_git_ops_redaction.py @@ -6,6 +6,7 @@ import pytest +from forge.integrations.source_control.contracts import GitCredentials from forge.workspace.git_ops import GitError, GitOperations from forge.workspace.manager import Workspace @@ -20,8 +21,9 @@ def _git_ops(tmp_path: Path) -> GitOperations: branch_name="forge/test", ticket_key="TEST-1", ) + credentials = GitCredentials(host="github.com", token=token) with patch("forge.workspace.git_ops.get_settings", return_value=settings): - return GitOperations(workspace) + return GitOperations(workspace, credentials) def test_clone_failure_redacts_token_from_git_error(tmp_path): @@ -52,9 +54,7 @@ def test_clone_failure_redacts_token_from_git_error(tmp_path): def test_git_error_constructor_redacts_tokens(): token = "gh" + "p_" + "abcdefghijklmnopqrstuvwxyz123456" - error = GitError( - f"remote: https://x-access-token:{token}@github.com/org/repo.git" - ) + error = GitError(f"remote: https://x-access-token:{token}@github.com/org/repo.git") assert "ghp_" not in str(error) assert "https://[REDACTED]@github.com/org/repo.git" in str(error) diff --git a/tests/unit/workspace/test_git_ops_single_branch.py b/tests/unit/workspace/test_git_ops_single_branch.py new file mode 100644 index 000000000..bd204541d --- /dev/null +++ b/tests/unit/workspace/test_git_ops_single_branch.py @@ -0,0 +1,132 @@ +"""Real-git tests for single-branch clones in direct mode. + +These reproduce the direct-mode failure where ``git clone --single-branch`` +restricts ``remote.origin.fetch`` to the default branch, so a later +unparameterized ``git fetch origin`` never creates the ``origin/`` +tracking ref that ``checkout_branch``/``pull_rebase`` rely on. Mocking +``_run_git`` cannot catch this because the bug lives in git's refspec +behavior, not in Forge's argument strings. +""" + +import subprocess +from pathlib import Path +from unittest.mock import MagicMock, patch + +from forge.integrations.source_control.contracts import GitCredentials +from forge.workspace.git_ops import GitOperations +from forge.workspace.manager import Workspace + +_GIT_ENV = { + "GIT_AUTHOR_NAME": "Test", + "GIT_AUTHOR_EMAIL": "test@example.com", + "GIT_COMMITTER_NAME": "Test", + "GIT_COMMITTER_EMAIL": "test@example.com", +} + + +def _run(cwd: Path, *args: str) -> None: + import os + + subprocess.run( + ["git", *args], + cwd=cwd, + check=True, + capture_output=True, + text=True, + env={**os.environ, **_GIT_ENV}, + ) + + +def _current_branch(repo: Path) -> str: + result = subprocess.run( + ["git", "rev-parse", "--abbrev-ref", "HEAD"], + cwd=repo, + check=True, + capture_output=True, + text=True, + ) + return result.stdout.strip() + + +def _make_remote(tmp_path: Path) -> tuple[Path, Path]: + """Create a bare remote holding a single ``main`` branch. + + Returns (bare_remote_path, seed_working_clone) so the caller can push + additional branches to the remote after the workspace has been cloned. + """ + seed = tmp_path / "seed" + seed.mkdir() + _run(seed, "init", "-b", "main") + (seed / "README.md").write_text("hello\n") + _run(seed, "add", ".") + _run(seed, "commit", "-m", "initial") + + remote = tmp_path / "remote.git" + _run(seed, "init", "--bare", str(remote)) + _run(seed, "remote", "add", "origin", str(remote)) + _run(seed, "push", "origin", "main") + return remote, seed + + +def _git_ops(repo_path: Path, branch: str) -> GitOperations: + workspace = Workspace( + path=repo_path, + repo_name="org/repo", + branch_name=branch, + ticket_key="TEST-123", + ) + credentials = GitCredentials(host="github.com", token="test-token") + with patch("forge.workspace.git_ops.get_settings", return_value=MagicMock()): + return GitOperations(workspace, credentials) + + +def test_checkout_branch_fetches_branch_created_after_single_branch_clone(tmp_path): + """A branch pushed to origin after a single-branch clone must still check out.""" + remote, seed = _make_remote(tmp_path) + repo_path = tmp_path / "repo" + git = _git_ops(repo_path, branch="forge/aisos-2420") + git.clone(repo_url=str(remote)) + _run(repo_path, "config", "user.name", "Test") + _run(repo_path, "config", "user.email", "test@example.com") + + # The feature branch appears on origin only after the single-branch clone, + # so its origin/ tracking ref does not exist locally yet. + _run(seed, "checkout", "-b", "forge/aisos-2420") + (seed / "feature.txt").write_text("feature\n") + _run(seed, "add", ".") + _run(seed, "commit", "-m", "feature commit") + _run(seed, "push", "origin", "forge/aisos-2420") + + git.checkout_branch("forge/aisos-2420", remote="origin") + + assert _current_branch(repo_path) == "forge/aisos-2420" + assert (repo_path / "feature.txt").exists() + + +def test_pull_rebase_direct_mode_rebases_after_single_branch_clone(tmp_path): + """pull_rebase(origin) rebases onto a remote branch a single-branch clone missed.""" + remote, seed = _make_remote(tmp_path) + repo_path = tmp_path / "repo" + git = _git_ops(repo_path, branch="forge/aisos-2420") + git.clone(repo_url=str(remote)) + _run(repo_path, "config", "user.name", "Test") + _run(repo_path, "config", "user.email", "test@example.com") + + # Local feature branch with an unpushed commit. + _run(repo_path, "checkout", "-b", "forge/aisos-2420") + (repo_path / "local.txt").write_text("local\n") + _run(repo_path, "add", ".") + _run(repo_path, "commit", "-m", "local commit") + + # Remote feature branch advances independently. + _run(seed, "checkout", "-b", "forge/aisos-2420") + (seed / "remote.txt").write_text("remote\n") + _run(seed, "add", ".") + _run(seed, "commit", "-m", "remote commit") + _run(seed, "push", "origin", "forge/aisos-2420") + + git.pull_rebase(remote="origin") + + # Rebase replayed the local commit on top of the fetched remote commit. + assert (repo_path / "remote.txt").exists() + assert (repo_path / "local.txt").exists() diff --git a/tests/unit/workspace/test_git_ops_sync.py b/tests/unit/workspace/test_git_ops_sync.py index 27815b94a..4c9e12b03 100644 --- a/tests/unit/workspace/test_git_ops_sync.py +++ b/tests/unit/workspace/test_git_ops_sync.py @@ -3,6 +3,7 @@ from pathlib import Path from unittest.mock import MagicMock, call, patch +from forge.integrations.source_control.contracts import GitCredentials from forge.workspace.git_ops import GitOperations from forge.workspace.manager import Workspace @@ -14,8 +15,9 @@ def _git_ops(tmp_path: Path) -> GitOperations: branch_name="forge/test-123", ticket_key="TEST-123", ) + credentials = GitCredentials(host="github.com", token="test-token") with patch("forge.workspace.git_ops.get_settings", return_value=MagicMock()): - return GitOperations(workspace) + return GitOperations(workspace, credentials) def test_pull_rebase_skips_missing_remote_branch_before_first_push(tmp_path): @@ -28,7 +30,11 @@ def test_pull_rebase_skips_missing_remote_branch_before_first_push(tmp_path): ): git.pull_rebase(remote="origin") - run_git.assert_called_once_with("fetch", "origin") + # The branch is fetched into its explicit remote-tracking ref so a + # --single-branch clone's narrow refspec cannot leave origin/ stale. + run_git.assert_called_once_with( + "fetch", "origin", "forge/test-123:refs/remotes/origin/forge/test-123", check=False + ) branch_exists.assert_called_once_with("forge/test-123", remote="origin") @@ -43,7 +49,7 @@ def test_pull_rebase_rebases_when_remote_branch_exists(tmp_path): git.pull_rebase(remote="fork") assert run_git.call_args_list == [ - call("fetch", "fork"), + call("fetch", "fork", "forge/test-123:refs/remotes/fork/forge/test-123", check=False), call("rebase", "fork/forge/test-123"), ] branch_exists.assert_called_once_with("forge/test-123", remote="fork") diff --git a/tests/workflow/test_task_takeover_graph.py b/tests/workflow/test_task_takeover_graph.py index aaecd673b..eee3fd334 100644 --- a/tests/workflow/test_task_takeover_graph.py +++ b/tests/workflow/test_task_takeover_graph.py @@ -94,9 +94,10 @@ class TestPathTransitions: ("qualitative_review", "run_qualitative_review"), ("create_pr", "create_pr"), ("teardown_workspace", "teardown_workspace"), - # Removed lifecycle nodes safely restart at triage rather than - # resolving to a node that no longer exists. - ("wait_for_ci_gate", "triage_check"), + # wait_for_ci_gate was merged into human_review_gate; a + # compatibility alias resumes checkpoints parked there instead of + # restarting the whole workflow from triage. + ("wait_for_ci_gate", "human_review_gate"), ("ci_evaluator", "ci_evaluator"), ("attempt_ci_fix", "ci_evaluator"), ("human_review_gate", "human_review_gate"),