diff --git a/.gitignore b/.gitignore index f4906988..8f9e9992 100644 --- a/.gitignore +++ b/.gitignore @@ -7,6 +7,7 @@ specs/ docs/github-app-setup.md # Python +.mypy_cache/ __pycache__/ *.py[cod] *$py.class diff --git a/CLAUDE.md b/CLAUDE.md index 3f852ad7..04401076 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -120,11 +120,16 @@ podman rm $(podman ps -a --filter name=forge- -q) ## Jira Comment Syntax -| Prefix | Effect | -|--------|--------| -| `!` | Revision request — triggers regeneration with feedback | +| Prefix / Command | Effect | +|------------------|--------| +| `!` | Revision request — triggers regeneration/revision with feedback | | `?` or `@forge ask` | Question — triggers Q&A answer | | `>option N` | RCA option selection (RCA Option Gate only) | +| `/forge approve` | Approve draft (Epic plan or Tasks) to provision tickets | +| `/forge remove ID` | Remove proposed draft item by ID and re-sequence | +| `/forge exclude ID` | Toggle exclusion of proposed draft item by ID | +| `/forge update ID key=val` | Update fields (`summary`, `description`, `repo`) of proposed draft item | +| `/forge add key=val` | Add a new proposed item to draft | | _(no prefix)_ | Informational — workflow ignores it | ## GitHub PR Comment Commands diff --git a/docs/developer-guide.md b/docs/developer-guide.md index 9bf63a0e..e9ab7eb8 100644 --- a/docs/developer-guide.md +++ b/docs/developer-guide.md @@ -858,6 +858,18 @@ curl -X POST http://localhost:8000/api/v1/webhooks/github \ | `/forge skip-gate ` | Skip named CI check | CI stages | | `/forge unskip-gate ` | Remove a skip | CI stages | +### Jira comment commands + +These commands are used on the parent Jira ticket during the draft review stages (Epic Plan and Tasks). + +| Command | Effect | Active at | +|---------|--------|-----------| +| `/forge approve` | Approve draft, provision sub-tickets, and delete draft attachment | `plan_approval_gate`, `task_approval_gate` | +| `/forge remove ` | Remove a draft item by local sequential ID | `plan_approval_gate`, `task_approval_gate` | +| `/forge exclude ` | Toggle the exclusion flag of a draft item | `plan_approval_gate`, `task_approval_gate` | +| `/forge update key=val` | Update fields (`summary`, `description`, `repo`) of a draft item | `plan_approval_gate`, `task_approval_gate` | +| `/forge add key=val` | Add a new proposed item to the draft | `plan_approval_gate`, `task_approval_gate` | + ### Jira labels | Label | Meaning | diff --git a/docs/guide/feature-workflow.md b/docs/guide/feature-workflow.md index 8abffdac..6e24c6cf 100644 --- a/docs/guide/feature-workflow.md +++ b/docs/guide/feature-workflow.md @@ -60,22 +60,29 @@ Forge generates a behavioral specification from the approved PRD, typically usin Forge breaks the feature into logical epics — high-level areas of work that map to implementation phases. -**Human action:** Review the epic plan. You have four options at this stage: +By default, Forge uses an interactive **Draft Review Flow** at this stage (unless YOLO mode is active): +1. Instead of creating Jira tickets immediately, Forge serializes the proposed epics into `forge-stories-draft.json` and uploads it as an attachment on the Feature ticket. +2. Forge posts a markdown table comment on the Feature ticket outlining the proposed Epics. +3. The workflow pauses at `plan_approval_gate`. -| Action | How | -|--------|-----| -| Approve | Change label to `forge:plan-approved` | -| Ask a question | Comment with `?` prefix — Forge answers without re-decomposing | -| Revise one epic | `!` comment on the **specific epic sub-ticket** — Forge updates only that epic | -| Redo the full decomposition | `!` comment on the **feature ticket** — Forge regenerates all epics with your feedback | +**Human action:** Review the epic plan draft. You have several options at this stage: + +| Action | How | Description | +|--------|-----|-------------| +| **Approve** | Comment `/forge approve` OR set label to `forge:plan-approved` | Forge provisions the Epic sub-tickets on Jira from the draft, deletes the draft attachment, and advances to Task Generation. | +| **Direct Edit** | Use `/forge` commands (e.g. `/forge update`, `/forge remove`, etc.) | Directly modify the draft attachment and regenerate the proposal comment. See [Jira Labels & Comments](labels.md) for a list of commands. | +| **Ask a question** | Comment with `?` prefix or `@forge ask` | Forge answers your question without regenerating the draft. | +| **Request revisions** | Comment with `!` prefix followed by your feedback | Forge uses LLM assistance to revise the entire draft JSON and update the proposal comment with your feedback. | + +If `forge:yolo` mode is active, the draft review is bypassed. Epics are created in Jira immediately, and the workflow automatically proceeds to Task Generation. ```mermaid flowchart TD Gate([plan_approval_gate]) - Gate -->|forge:plan-approved| Next[Generate Tasks] + Gate -->|forge:plan-approved or /forge approve| Next[Generate Tasks] Gate -->|"? on feature ticket"| QA[Answer Question] Gate -->|"! on feature ticket"| Regen[Regenerate All Epics] - Gate -->|"! on epic sub-ticket"| Update[Update Single Epic] + Gate -->|"/forge update/remove/exclude/add"| Update[Modify Draft] QA --> Gate Regen --> Gate Update --> Gate @@ -87,22 +94,29 @@ flowchart TD Forge generates granular implementation tasks scoped to individual repositories. Each task is sized to fit in a single container execution pass. -**Human action:** Review the tasks. You have four options at this stage: +By default, Forge uses an interactive **Draft Review Flow** at this stage (unless YOLO mode is active): +1. Instead of creating Jira tickets immediately, Forge serializes the proposed tasks into `forge-tasks-draft.json` and uploads it as an attachment on the Feature ticket. +2. Forge posts a markdown table comment on the Feature ticket outlining the proposed Tasks. +3. The workflow pauses at `task_approval_gate`. + +**Human action:** Review the task draft. You have several options at this stage: + +| Action | How | Description | +|--------|-----|-------------| +| **Approve** | Comment `/forge approve` OR set label to `forge:task-approved` | Forge provisions the Task sub-tickets on Jira from the draft, deletes the draft attachment, and advances to Implementation. | +| **Direct Edit** | Use `/forge` commands (e.g. `/forge update`, `/forge remove`, etc.) | Directly modify the draft attachment and regenerate the proposal comment. See [Jira Labels & Comments](labels.md) for a list of commands. | +| **Ask a question** | Comment with `?` prefix or `@forge ask` | Forge answers your question without regenerating the draft. | +| **Request revisions** | Comment with `!` prefix followed by your feedback | Forge uses LLM assistance to revise the entire draft JSON and update the proposal comment with your feedback. | -| Action | How | -|--------|-----| -| Approve | Change label to `forge:task-approved` | -| Ask a question | Comment with `?` prefix — Forge answers without regenerating | -| Revise one task | `!` comment on the **specific task sub-ticket** — Forge updates only that task | -| Regenerate all tasks | `!` comment on the **feature or epic ticket** — Forge regenerates the full task list with your feedback | +If `forge:yolo` mode is active, the draft review is bypassed. Tasks are created in Jira immediately, and the workflow automatically proceeds to Implementation. ```mermaid flowchart TD Gate([task_approval_gate]) - Gate -->|forge:task-approved| Next[Implement Tasks] + Gate -->|forge:task-approved or /forge approve| Next[Implement Tasks] Gate -->|"? on ticket"| QA[Answer Question] Gate -->|"! on feature/epic"| Regen[Regenerate All Tasks] - Gate -->|"! on task sub-ticket"| Update[Update Single Task] + Gate -->|"/forge update/remove/exclude/add"| Update[Modify Draft] QA --> Gate Regen --> Gate Update --> Gate diff --git a/docs/guide/labels.md b/docs/guide/labels.md index 8cc74f8a..eb0f2c99 100644 --- a/docs/guide/labels.md +++ b/docs/guide/labels.md @@ -51,13 +51,33 @@ Standalone Tasks and Epics can be processed with the standard `forge:managed` la **Starting a workflow:** Create a Jira issue and add `forge:managed`. Forge detects the issue type and begins the appropriate pipeline: Feature/Story, Bug, or standalone Task/Epic takeover. -**Approving a stage:** When Forge posts a PRD, spec, or other artifact, it sets the `forge:*-pending` label. Change it to `forge:*-approved` to advance the workflow. Do not add the approved label manually before Forge posts — it won't be recognized until the pending state is set. +**Approving a stage:** When Forge posts an artifact (such as a PRD or Spec), it sets the `forge:*-pending` label. You can approve it by changing the label to `forge:*-approved` to advance the workflow. For draft-based stages (Epic Plan and Tasks), you can also approve by commenting `/forge approve` on the ticket. -**Requesting revisions:** Start a comment with `!` followed by your feedback. Forge regenerates the artifact and resets the pending label. +**Interactive Draft Review:** For Epic Decomposition and Task Generation stages, Forge uses a draft-based review flow by default (unless `forge:yolo` mode is active). +1. Instead of creating sub-tickets immediately, Forge serializes the proposed items into a JSON draft file (`forge-stories-draft.json` or `forge-tasks-draft.json`) and uploads it as a Jira attachment. +2. Forge posts a formatted markdown table comment on the ticket detailing the proposed plan. +3. While the stage is pending, you can modify the draft directly using **Jira comment commands** (see below) or request a natural language revision. +4. Once you approve (via `/forge approve` or setting the approved label), Forge downloads the draft, provisions the actual Jira tickets from it, and deletes the draft attachment. -**Asking questions:** Start a comment with `?` or `@forge ask`. Forge answers without advancing or regenerating. +### Jira Comment Commands -**Informational comments:** Comments without a recognized prefix (`!`, `?`, `@forge ask`, `>option`) are ignored by the workflow — use them for team discussion without triggering Forge. +For stages using the draft-based review flow (Epic Plan and Tasks), you can post comments on the parent ticket with the following commands: + +| Command | Description | Example | +|---------|-------------|---------| +| `/forge approve` | Approve the draft, provision all non-excluded items as Jira tickets, and delete the draft attachment. | `/forge approve` | +| `/forge remove ` | Remove a draft item by its local sequential ID. Remaining items are automatically re-sequenced. | `/forge remove 3` | +| `/forge exclude ` | Toggle the exclusion flag of a draft item. Excluded items are skipped during ticket provisioning. | `/forge exclude 2` | +| `/forge update key=val` | Update fields of a draft item (supported keys: `summary`, `description`, `repo`). | `/forge update 1 repo="my-org/custom-repo"` | +| `/forge add key=val` | Add a new proposed item to the draft. | `/forge add summary="New Story" repo="my-org/repo"` | + +*Note: Successful command/revision comments are automatically edited by Forge to prepend `✅`. If a command or revision fails, Forge posts a comment detailing the error with a leading `❌`.* + +**Requesting revisions:** Start a comment with `!` followed by your feedback (e.g., `! update the repositories to use the new service`). For standard artifacts, Forge regenerates them. For drafts, Forge uses LLM assistance to revise the draft JSON attachment and update the proposed plan table. + +**Asking questions:** Start a comment with `?` or `@forge ask`. Forge answers without advancing or regenerating/modifying the drafts. + +**Informational comments:** Comments without a recognized prefix (such as `!`, `?`, `@forge ask`, `>option`, or `/forge`) are ignored by the workflow — use them for team discussion without triggering Forge. **Handling failures:** When `forge:blocked` appears, read the Forge comment for the error. Fix the underlying issue if needed, then add `forge:retry`. diff --git a/src/forge/config.py b/src/forge/config.py index 26a61b12..c6813f1c 100644 --- a/src/forge/config.py +++ b/src/forge/config.py @@ -339,6 +339,11 @@ def ignored_ci_checks(self) -> list[str]: default=0.5, description="Webhook acknowledgment timeout in seconds" ) + yolo_mode: bool = Field( + default=False, + description="Autonomous mode - skip all artifact approval gates", + ) + # Container Configuration container_image: str = Field( default="localhost/forge-dev:latest", diff --git a/src/forge/integrations/agents/agent.py b/src/forge/integrations/agents/agent.py index 974330c7..5ca5e01b 100644 --- a/src/forge/integrations/agents/agent.py +++ b/src/forge/integrations/agents/agent.py @@ -1272,6 +1272,89 @@ async def answer_question( logger.info(f"Generated answer ({len(result)} chars)") return result.strip() if result else "" + async def revise_draft_with_feedback( + self, + draft_content: str, + feedback: str, + context: dict[str, Any] | None = None, + ) -> str: + """Revise draft content based on user feedback. + + Uses the 'revision-draft' prompt template to guide the LLM to output + the revised draft JSON. + + Args: + draft_content: The current draft JSON content. + feedback: Natural language feedback. + context: Optional context from the workflow state. + + Returns: + The updated draft JSON string. + """ + from langchain_core.output_parsers import StrOutputParser + + # Format context into a readable string/JSON + context_str = json.dumps(context, indent=2) if context else "None provided" + + # Load the prompt template using project's load_prompt + prompt_text = load_prompt( + "revision-draft", + draft_content=draft_content, + feedback=feedback, + context=context_str, + ) + + model = self._create_model() + chain = model | StrOutputParser() + + logger.info("Revising draft using direct LangChain model chain") + response = await chain.ainvoke(prompt_text) + + # Strip preamble/narration and validate as JSON + cleaned_text = response.strip() + + # Check markdown code blocks first + pattern = r"```(?:json)?\s*([\s\S]*?)\s*```" + match = re.search(pattern, cleaned_text) + if match: + cleaned_text = match.group(1).strip() + else: + # If no code block, look for the JSON object/list boundary + # Find the first occurrence of '{' or '[' and the last of '}' or ']' + start_brace = cleaned_text.find("{") + start_bracket = cleaned_text.find("[") + + # Determine which starts first + start_idx = -1 + if start_brace != -1 and start_bracket != -1: + start_idx = min(start_brace, start_bracket) + elif start_brace != -1: + start_idx = start_brace + elif start_bracket != -1: + start_idx = start_bracket + + if start_idx != -1: + # Find the last brace or bracket + end_brace = cleaned_text.rfind("}") + end_bracket = cleaned_text.rfind("]") + end_idx = max(end_brace, end_bracket) + + if end_idx > start_idx: + cleaned_text = cleaned_text[start_idx : end_idx + 1].strip() + + try: + parsed_json = json.loads(cleaned_text) + validated_json_str = json.dumps(parsed_json, indent=2) + logger.info( + f"Successfully revised draft and validated JSON ({len(validated_json_str)} chars)" + ) + return validated_json_str + except json.JSONDecodeError as e: + logger.error(f"Failed to parse LLM response as valid JSON: {e}\nResponse: {response}") + raise ValueError( + f"Failed to parse LLM response as valid JSON: {e}\nResponse: {response}" + ) + async def close(self) -> None: """Close the agent and cleanup resources.""" pass diff --git a/src/forge/integrations/jira/client.py b/src/forge/integrations/jira/client.py index b519e58c..ea1d3c6a 100644 --- a/src/forge/integrations/jira/client.py +++ b/src/forge/integrations/jira/client.py @@ -81,7 +81,6 @@ async def _get_client(self) -> httpx.AsyncClient: ), headers={ "Accept": "application/json", - "Content-Type": "application/json", }, timeout=30.0, ) @@ -415,58 +414,76 @@ async def add_attachment( Args: issue_key: The Jira issue key. filename: Name for the attachment file. - content: File content (string or bytes). - content_type: MIME type of the content. + content: File content as string or bytes. + content_type: The content type of the file. Returns: The attachment metadata from Jira API. """ - # Attachments require a separate client without JSON content-type - async with httpx.AsyncClient( - base_url=self.base_url, - auth=( - self.settings.jira_user_email, - self.settings.jira_api_token.get_secret_value(), - ), - headers={ - "Accept": "application/json", - "X-Atlassian-Token": "no-check", # Required for attachments - }, - timeout=60.0, - ) as client: - # Convert string to bytes if needed - if isinstance(content, str): - content = content.encode("utf-8") - - files = {"file": (filename, content, content_type)} - response = await client.post( - f"/issue/{issue_key}/attachments", - files=files, - ) - response.raise_for_status() - data = response.json() - logger.info(f"Added attachment {filename} to {issue_key}") - return data[0] if data else {} + if isinstance(content, str): + content = content.encode("utf-8") + + if content_type == "text/markdown" and filename.endswith(".json"): + content_type = "application/json" + + headers = { + "X-Atlassian-Token": "no-check", + } + files = {"file": (filename, content, content_type)} + + response = await self._request_with_retry( + "POST", + f"/issue/{issue_key}/attachments", + headers=headers, + files=files, + ) + response.raise_for_status() + data = response.json() + logger.info(f"Added attachment {filename} to {issue_key}") + return data[0] if data else {} async def get_attachments(self, issue_key: str) -> list[dict[str, Any]]: - """Get all attachments for a Jira issue. + """Get all attachments for a Jira issue by querying the issue's details. Args: issue_key: The Jira issue key. Returns: - List of attachment metadata dicts with 'id', 'filename', 'size', etc. + A list of attachment metadata dicts containing id, filename, and content URL. """ - client = await self._get_client() - response = await client.get( + response = await self._request_with_retry( + "GET", f"/issue/{issue_key}", params={"fields": "attachment"}, ) response.raise_for_status() data = response.json() attachments = data.get("fields", {}).get("attachment", []) - logger.debug(f"Found {len(attachments)} attachments on {issue_key}") - return attachments + + result = [] + for att in attachments: + result.append( + { + "id": att.get("id"), + "filename": att.get("filename"), + "content_url": att.get("content"), + } + ) + logger.debug(f"Found {len(result)} attachments on {issue_key}") + return result + + async def download_attachment(self, content_url: str) -> bytes: + """Download attachment raw binary content from the given content URL. + + Args: + content_url: The full URL to download the attachment. + + Returns: + The raw binary content of the attachment. + """ + response = await self._request_with_retry("GET", content_url) + response.raise_for_status() + return response.content async def delete_attachment(self, attachment_id: str) -> None: """Delete an attachment by ID. @@ -474,8 +491,7 @@ async def delete_attachment(self, attachment_id: str) -> None: Args: attachment_id: The Jira attachment ID. """ - client = await self._get_client() - response = await client.delete(f"/attachment/{attachment_id}") + response = await self._request_with_retry("DELETE", f"/attachment/{attachment_id}") response.raise_for_status() logger.info(f"Deleted attachment {attachment_id}") @@ -599,6 +615,30 @@ async def add_comment(self, issue_key: str, body: str) -> JiraComment: logger.info(f"Added comment to {issue_key}") return JiraComment.from_api_response(data) + async def edit_comment(self, issue_key: str, comment_id: str, body: str) -> JiraComment: + """Edit an existing comment on a Jira issue. + + Args: + issue_key: The Jira issue key. + comment_id: The Jira comment ID. + body: New comment text content. + + Returns: + The updated JiraComment. + """ + adf_content = self._text_to_adf(body) + + # Edit the comment using request_with_retry to handle transient rate limits robustly + response = await self._request_with_retry( + "PUT", + f"/issue/{issue_key}/comment/{comment_id}", + json={"body": adf_content}, + ) + response.raise_for_status() + data = response.json() + logger.info(f"Edited comment {comment_id} on {issue_key}") + return JiraComment.from_api_response(data) + async def add_error_comment( self, issue_key: str, diff --git a/src/forge/models/__init__.py b/src/forge/models/__init__.py index 17a8b697..d67d6881 100644 --- a/src/forge/models/__init__.py +++ b/src/forge/models/__init__.py @@ -1,6 +1,7 @@ """Domain models for Forge orchestrator.""" from forge.models.artifacts import Epic, Feature, Task +from forge.models.draft import DraftItem, ForgeDecompositionDraft from forge.models.events import EventSource, EventStatus, WebhookEvent from forge.models.workflow import ( ForgeLabel, @@ -23,6 +24,8 @@ "Feature", "Epic", "Task", + "DraftItem", + "ForgeDecompositionDraft", # Event models "WebhookEvent", "EventSource", diff --git a/src/forge/models/draft.py b/src/forge/models/draft.py new file mode 100644 index 00000000..3b17ae96 --- /dev/null +++ b/src/forge/models/draft.py @@ -0,0 +1,77 @@ +"""Data models for decomposing draft artifacts.""" + +from datetime import datetime +from typing import Literal + +from pydantic import BaseModel, model_validator + + +class DraftItem(BaseModel): + """Represents an individual proposed Story or Task inside a draft.""" + + model_config = {"extra": "forbid"} + + id: int + """Local sequential ID, e.g., 1, 2, 3.""" + + summary: str + """Brief summary of the proposed item.""" + + description: str + """Detailed description of the proposed item.""" + + repo: str + """Target repository name.""" + + acceptance_criteria: list[str] + """List of acceptance criteria for this item.""" + + excluded: bool = False + """Whether this item should be excluded from ticket creation.""" + + epic_key: str | None = None + """Optional Jira key of the parent epic for this task.""" + + +class ForgeDecompositionDraft(BaseModel): + """Represents the wrapper of all draft items and execution metadata.""" + + parent_key: str + """Jira key of the parent feature or epic.""" + + phase: Literal["stories", "tasks"] + """Phase of the draft, either "stories" or "tasks".""" + + items: list[DraftItem] + """List of draft items (proposed Stories or Tasks).""" + + version: int = 1 + """Draft schema version.""" + + created_at: datetime + """Timestamp when this draft was created.""" + + updated_at: datetime + """Timestamp when this draft was last updated.""" + + @model_validator(mode="after") + def _validate_sequential_ids(self) -> "ForgeDecompositionDraft": + """Validate that local item IDs are unique and sequential (1, 2, 3...) within the draft.""" + if not self.items: + return self + + ids = [item.id for item in self.items] + + # Check uniqueness + if len(ids) != len(set(ids)): + raise ValueError("Draft item IDs must be unique.") + + # Check that IDs are sequential starting from 1 + sorted_ids = sorted(ids) + expected_ids = list(range(1, len(self.items) + 1)) + if sorted_ids != expected_ids: + raise ValueError( + f"Draft item IDs must be sequential starting from 1. Got: {sorted_ids}, expected: {expected_ids}" + ) + + return self diff --git a/src/forge/orchestrator/worker.py b/src/forge/orchestrator/worker.py index a0746b65..e5136665 100644 --- a/src/forge/orchestrator/worker.py +++ b/src/forge/orchestrator/worker.py @@ -9,8 +9,9 @@ import sys import uuid from dataclasses import replace as dataclass_replace +from datetime import UTC, datetime from pathlib import Path -from typing import Any +from typing import Any, cast from forge.api.routes.metrics import ( record_workflow_completed, @@ -18,8 +19,10 @@ record_workflow_started, ) 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.models.draft import DraftItem, ForgeDecompositionDraft 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 @@ -28,6 +31,8 @@ from forge.skills.orchestrator import ensure_skills from forge.skills.utils import extract_project_key from forge.utils.redaction import redact_secrets +from forge.workflow.gates.plan_approval import provision_epics_from_draft +from forge.workflow.gates.task_approval import provision_tasks_from_draft from forge.workflow.nodes.error_handler import notify_error from forge.workflow.pr_state import ( activate_pull_request_for_event, @@ -42,7 +47,16 @@ is_bot_sender, triage_automated_review, ) -from forge.workflow.utils.comment_classifier import CommentType, classify_comment +from forge.workflow.utils.comment_classifier import ( + CommentType, + classify_comment, + parse_comment_command, +) +from forge.workflow.utils.draft_manager import ( + FORGE_STORIES_DRAFT_FILENAME, + FORGE_TASKS_DRAFT_FILENAME, + DraftManager, +) from forge.workflow.utils.jira_status import post_status_comment from forge.workflow.utils.proposal_review_threads import ( reply_to_proposal_decisions, @@ -109,6 +123,12 @@ async def _report_new_workflow_error(result: dict, error_before_invoke: str | No "rca_option_gate", } +_PENDING_APPROVAL_GATES = { + "plan_approval_gate", + "task_plan_approval_gate", + "task_approval_gate", +} + class OrchestratorWorker: """Worker that processes workflow events from Redis queue.""" @@ -705,6 +725,24 @@ async def _handle_resume_event( is_retry = True logger.info(f"Detected retry signal via forge:retry label for {current_node}") + # Direct check for forge:plan-approved and forge:task-approved additions + to_lower = to_labels.lower() if to_labels else "" + from_lower = from_labels.lower() if from_labels else "" + if ( + "forge:plan-approved" in to_lower + and "forge:plan-approved" not in from_lower + and current_node in ("plan_approval_gate", "task_plan_approval_gate") + ): + is_approved = True + logger.info(f"Detected forge:plan-approved addition on {message.ticket_key}") + if ( + "forge:task-approved" in to_lower + and "forge:task-approved" not in from_lower + and current_node == "task_approval_gate" + ): + is_approved = True + logger.info(f"Detected forge:task-approved addition on {message.ticket_key}") + # Check for approval labels - but only if it matches the current stage if "approved" in to_labels.lower() and "pending" in from_labels.lower(): # Validate the approval matches the workflow stage @@ -790,6 +828,170 @@ async def _handle_resume_event( comment_body = self._extract_text_from_adf(comment_body) if comment_body.strip(): + # Check for interactive comment commands or natural language feedback when paused in PENDING_APPROVAL (BR-006) + if current_state.get("is_paused") and current_node in _PENDING_APPROVAL_GATES: + parsed_cmd = parse_comment_command(comment_body) + is_forge_cmd = parsed_cmd is not None + # Revision comment check allows leading whitespace for consistency with the comment classifier + is_revision_comment = bool(re.match(r"^\s*!", comment_body)) + + if is_forge_cmd or is_revision_comment: + filename = ( + FORGE_STORIES_DRAFT_FILENAME + if current_node in ("plan_approval_gate", "task_plan_approval_gate") + else FORGE_TASKS_DRAFT_FILENAME + ) + + jira = JiraClient() + has_draft = False + try: + attachments = await jira.get_attachments(message.ticket_key) + has_draft = any(att.get("filename") == filename for att in attachments) + except Exception as list_err: + logger.warning(f"Could not list attachments to check draft: {list_err}") + + if is_forge_cmd or (is_revision_comment and has_draft): + original_draft = None + try: + try: + original_draft = await DraftManager.get_draft_attachment( + jira, message.ticket_key, filename + ) + except Exception as get_err: + logger.warning( + f"Could not download original draft for rollback: {get_err}" + ) + + if is_forge_cmd and parsed_cmd is not None: + if parsed_cmd.get("command") == "approve": + is_approved = True + if current_node in ( + "plan_approval_gate", + "task_plan_approval_gate", + ): + await jira.set_workflow_label( + message.ticket_key, ForgeLabel.PLAN_APPROVED + ) + elif current_node == "task_approval_gate": + await jira.set_workflow_label( + message.ticket_key, ForgeLabel.TASK_APPROVED + ) + else: + if not original_draft: + raise ValueError( + f"Draft attachment '{filename}' not found for modification." + ) + + draft_json = [ + item.model_dump() for item in original_draft.items + ] + mutated_json = DraftManager.apply_draft_modification( + draft_json, parsed_cmd + ) + + updated_items = [ + DraftItem.model_validate(item) for item in mutated_json + ] + updated_draft = ForgeDecompositionDraft( + parent_key=original_draft.parent_key, + phase=original_draft.phase, + items=updated_items, + version=original_draft.version, + created_at=original_draft.created_at, + updated_at=datetime.now(UTC), + ) + + await DraftManager.save_draft_attachment( + jira, message.ticket_key, updated_draft, filename + ) + + # Update the original review comment with the new breakdown (SC-002) + await self._update_original_review_comment( + jira, message.ticket_key, updated_draft + ) + + comment_id = comment.get("id") if comment else None + if comment_id: + await jira.edit_comment( + message.ticket_key, + comment_id, + f"✅ {comment_body}", + ) + + elif is_revision_comment: + if not original_draft: + raise ValueError( + f"Draft attachment '{filename}' not found for revision." + ) + + feedback_text = re.sub(r"^\s*!\s*", "", comment_body) + if not feedback_text: + raise ValueError("Revision feedback cannot be empty.") + + agent = ForgeAgent() + try: + revised_json_str = await agent.revise_draft_with_feedback( + draft_content=original_draft.model_dump_json(), + feedback=feedback_text, + context={ + "ticket_key": message.ticket_key, + "current_node": current_node, + }, + ) + finally: + await agent.close() + + updated_draft = ForgeDecompositionDraft.model_validate_json( + revised_json_str + ) + + await DraftManager.save_draft_attachment( + jira, message.ticket_key, updated_draft, filename + ) + + # Update the original review comment with the new breakdown (SC-002) + await self._update_original_review_comment( + jira, message.ticket_key, updated_draft + ) + + comment_id = comment.get("id") if comment else None + if comment_id: + await jira.edit_comment( + message.ticket_key, comment_id, f"✅ {comment_body}" + ) + + if not parsed_cmd or parsed_cmd.get("command") != "approve": + await jira.close() + return current_state + + except Exception as e: + logger.error( + f"Failed to process comment command/revision: {e}", + exc_info=True, + ) + if original_draft: + try: + await DraftManager.save_draft_attachment( + jira, message.ticket_key, original_draft, filename + ) + except Exception as rollback_err: + logger.error( + f"Failed to roll back draft attachment: {rollback_err}", + exc_info=True, + ) + + error_comment_text = f"❌ Forge command/revision failed: {str(e)}" + try: + await jira.add_comment(message.ticket_key, error_comment_text) + except Exception as post_err: + logger.error( + f"Failed to post error comment: {post_err}", exc_info=True + ) + + await jira.close() + return current_state + await jira.close() + # >option N detection for rca_option_gate (runs before general classification) if current_node == "rca_option_gate": option_match = _OPTION_PATTERN.search(comment_body) @@ -1113,7 +1315,7 @@ async def _handle_resume_event( repo_full = payload.get("repository", {}).get("full_name", "") pr_number = payload.get("pull_request", {}).get("number") review_id = review.get("id") - inline_comments: list[dict[str, Any]] = [] + inline_comments = [] if repo_full and pr_number and review_id: _owner, _repo = repo_full.split("/", 1) gh = GitHubClient() @@ -1560,6 +1762,54 @@ async def _handle_resume_event( updated_state["automated_review_revision_count"] = 0 updated_state["automated_review_revision_pending"] = False updated_state["proposal_review_decisions"] = [] + + # Ticket provisioning step on approval! + # Note on split ownership: The worker call site handles unpausing from human manual comments (webhook triggers) + # to catch and report provisioning errors early without breaking the LangGraph execution flow. + if current_node == "plan_approval_gate" and not updated_state.get("epic_keys"): + jira = JiraClient() + try: + epic_keys = await provision_epics_from_draft(cast(Any, updated_state), jira) + updated_state["epic_keys"] = epic_keys + except Exception as e: + logger.error( + f"Failed ticket provisioning during plan approval for {message.ticket_key}: {e}", + exc_info=True, + ) + # Keep paused in PENDING_APPROVAL and post error comment + error_comment_text = f"❌ Ticket provisioning failed: {str(e)}" + try: + await jira.add_comment(message.ticket_key, error_comment_text) + except Exception as post_err: + logger.error(f"Failed to post error comment: {post_err}", exc_info=True) + return current_state + finally: + await jira.close() + + elif current_node == "task_approval_gate" and not updated_state.get("task_keys"): + # Note on split ownership: The worker call site handles unpausing from human manual comments (webhook triggers) + # to catch and report provisioning errors early without breaking the LangGraph execution flow. + jira = JiraClient() + try: + task_keys, tasks_by_repo = await provision_tasks_from_draft( + cast(Any, updated_state), jira + ) + updated_state["task_keys"] = task_keys + updated_state["tasks_by_repo"] = tasks_by_repo + except Exception as e: + logger.error( + f"Failed ticket provisioning during task approval for {message.ticket_key}: {e}", + exc_info=True, + ) + # Keep paused in PENDING_APPROVAL and post error comment + error_comment_text = f"❌ Ticket provisioning failed: {str(e)}" + try: + await jira.add_comment(message.ticket_key, error_comment_text) + except Exception as post_err: + logger.error(f"Failed to post error comment: {post_err}", exc_info=True) + return current_state + finally: + await jira.close() elif is_question: # Unpause so answer_question node runs, it will re-pause after answering updated_state["is_paused"] = False @@ -1853,6 +2103,39 @@ async def _post_rebase_feedback( except Exception as e: logger.warning(f"Failed to post rebase feedback: {e}") + async def _update_original_review_comment( + self, jira: JiraClient, ticket_key: str, draft: ForgeDecompositionDraft + ) -> None: + """Update the original review comment with the new breakdown (SC-002). + + Args: + jira: The Jira client. + ticket_key: The Jira ticket key. + draft: The updated draft decomposition. + """ + try: + comments_list = await jira.get_comments(ticket_key) + target_prefix = ( + "### 📋 Proposed Epics Draft" + if draft.phase == "stories" + else "### 📋 Proposed Tasks Draft" + ) + review_comment_id = None + for c in reversed(comments_list): + if c.body.startswith(target_prefix): + review_comment_id = c.id + break + if review_comment_id: + new_comment_body = DraftManager.format_review_comment(draft) + await jira.edit_comment( + ticket_key, + review_comment_id, + new_comment_body, + ) + logger.info(f"Edited original review comment {review_comment_id}") + except Exception as c_err: + logger.warning(f"Could not update original review comment: {c_err}") + async def _post_terminal_error_comment(self, ticket_key: str, error: str) -> None: """Post a comment explaining how to retry a terminal error. diff --git a/src/forge/prompts/v1/revision-draft.md b/src/forge/prompts/v1/revision-draft.md new file mode 100644 index 00000000..bcf5aec0 --- /dev/null +++ b/src/forge/prompts/v1/revision-draft.md @@ -0,0 +1,22 @@ +Please revise the following draft JSON list based on the parent issue context and natural language feedback. + +## Current Draft Content (JSON) + +{draft_content} + +## Parent Issue Context + +{context} + +## Feedback / Revision Request + +{feedback} + +## Instructions + +- Revise the draft JSON list to incorporate the feedback. +- Preserve the existing JSON structure and fields. +- Ensure all items in the draft maintain valid structure and formatting. +- Make sure the output is a valid JSON string representing the updated draft. +- You MUST output ONLY the raw JSON string. Do not include any preamble, introduction, explanation, or markdown code block syntax (like ```json ... ```). +- Start directly with the opening curly brace `{` or square bracket `[`. diff --git a/src/forge/workflow/gates/plan_approval.py b/src/forge/workflow/gates/plan_approval.py index cd7acf3e..bbd353dd 100644 --- a/src/forge/workflow/gates/plan_approval.py +++ b/src/forge/workflow/gates/plan_approval.py @@ -9,12 +9,16 @@ """ import logging +from typing import TYPE_CHECKING, Any, cast from langgraph.graph import END from forge.api.routes.metrics import record_approval, record_revision_requested from forge.workflow.feature.state import FeatureState as WorkflowState -from forge.workflow.utils import set_paused +from forge.workflow.utils import check_yolo_mode, set_paused + +if TYPE_CHECKING: + from forge.integrations.jira.client import JiraClient logger = logging.getLogger(__name__) @@ -37,8 +41,10 @@ def plan_approval_gate(state: WorkflowState) -> WorkflowState: epic_keys = state.get("epic_keys", []) epic_count = len(epic_keys) + is_yolo = check_yolo_mode(state) + # Validate that we actually have epics to approve - if epic_count == 0: + if epic_count == 0 and is_yolo: logger.error( f"Plan approval gate reached with 0 Epics for {ticket_key}. " "This indicates epic decomposition failed. Routing back to retry." @@ -52,10 +58,10 @@ def plan_approval_gate(state: WorkflowState) -> WorkflowState: logger.info(f"Plan approval gate: pausing workflow for {ticket_key} ({epic_count} Epics)") - return set_paused(state, "plan_approval_gate") + return cast(WorkflowState, set_paused(cast(dict[str, Any], state), "plan_approval_gate")) -def route_plan_approval(state: WorkflowState) -> str: +async def route_plan_approval(state: WorkflowState) -> str: """Route based on plan approval status. Args: @@ -70,7 +76,7 @@ def route_plan_approval(state: WorkflowState) -> str: return "answer_question" # YOLO mode: auto-approve without human input - if state.get("yolo_mode"): + if check_yolo_mode(state): logger.info(f"YOLO mode: auto-approving plan for {state['ticket_key']}") record_approval("plan") return "generate_tasks" @@ -99,7 +105,79 @@ def route_plan_approval(state: WorkflowState) -> str: ) return END + # Handle standard (non-YOLO) approval draft ticket provisioning + # Note on split ownership: The gate nodes handle YOLO/autonomous paths where no manual comment or webhook + # is processed (so we must provision here), while the worker handles manual/human comment webhook triggers. + if not state.get("epic_keys"): + from forge.integrations.jira.client import JiraClient + + jira = JiraClient() + try: + epic_keys = await provision_epics_from_draft(state, jira) + # Store the newly created keys + state["epic_keys"] = epic_keys + except Exception as e: + logger.error( + f"Failed ticket provisioning during plan approval for {state['ticket_key']}: {e}", + exc_info=True, + ) + raise + finally: + await jira.close() + # All Epics approved, proceed to task generation logger.info(f"Epics approved for {state['ticket_key']}, proceeding to task generation") record_approval("plan") return "generate_tasks" + + +async def provision_epics_from_draft(state: WorkflowState, jira: "JiraClient") -> list[str]: + """Provision Epics from the plan draft attachment on Jira. + + Args: + state: The workflow state dictionary. + jira: An active JiraClient instance. + + Returns: + List of created Epic ticket keys. + """ + ticket_key = state["ticket_key"] + from forge.models.workflow import ForgeLabel + from forge.workflow.utils.draft_manager import FORGE_STORIES_DRAFT_FILENAME, DraftManager + + logger.info(f"Downloading plan draft for {ticket_key}") + draft = await DraftManager.get_draft_attachment(jira, ticket_key, FORGE_STORIES_DRAFT_FILENAME) + if not draft: + raise ValueError(f"Approved draft {FORGE_STORIES_DRAFT_FILENAME} not found on {ticket_key}") + + parent_issue = await jira.get_issue(ticket_key) + project_key = parent_issue.project_key + + epic_keys = [] + for item in draft.items: + if item.excluded: + logger.info(f"Skipping excluded plan item {item.id}: {item.summary}") + continue + + labels = [ + ForgeLabel.FORGE_MANAGED.value, + f"forge:parent:{ticket_key}", + ] + if item.repo and "/" in item.repo: + labels.append(f"repo:{item.repo}") + + epic_key = await jira.create_epic( + project_key=project_key, + summary=item.summary, + description=item.description, + parent_key=ticket_key, + labels=labels, + ) + epic_keys.append(epic_key) + + # Delete the draft only after 100% successful ticket creation + await DraftManager.delete_draft_attachment(jira, ticket_key, FORGE_STORIES_DRAFT_FILENAME) + logger.info( + f"Successfully provisioned {len(epic_keys)} Epics for {ticket_key} and deleted draft" + ) + return epic_keys diff --git a/src/forge/workflow/gates/task_approval.py b/src/forge/workflow/gates/task_approval.py index 32daceab..b5c04a51 100644 --- a/src/forge/workflow/gates/task_approval.py +++ b/src/forge/workflow/gates/task_approval.py @@ -9,12 +9,16 @@ """ import logging +from typing import TYPE_CHECKING, Any, cast from langgraph.graph import END from forge.api.routes.metrics import record_approval, record_revision_requested from forge.workflow.feature.state import FeatureState as WorkflowState -from forge.workflow.utils import set_paused +from forge.workflow.utils import check_yolo_mode, set_paused + +if TYPE_CHECKING: + from forge.integrations.jira.client import JiraClient logger = logging.getLogger(__name__) @@ -41,8 +45,10 @@ def task_approval_gate(state: WorkflowState) -> WorkflowState: task_keys = state.get("task_keys", []) task_count = len(task_keys) + is_yolo = check_yolo_mode(state) + # Validate that we actually have tasks to approve - if task_count == 0: + if task_count == 0 and is_yolo: logger.error( f"Task approval gate reached with 0 Tasks for {ticket_key}. " "This indicates task generation failed. Routing back to retry." @@ -59,10 +65,10 @@ def task_approval_gate(state: WorkflowState) -> WorkflowState: f"({task_count} Tasks pending implementation approval)" ) - return set_paused(state, "task_approval_gate") + return cast(WorkflowState, set_paused(cast(dict[str, Any], state), "task_approval_gate")) -def route_task_approval(state: WorkflowState) -> str: +async def route_task_approval(state: WorkflowState) -> str: """Route based on task approval status. Routing logic: @@ -87,7 +93,7 @@ def route_task_approval(state: WorkflowState) -> str: return "answer_question" # YOLO mode: auto-approve without human input - if state.get("yolo_mode"): + if check_yolo_mode(state): logger.info(f"YOLO mode: auto-approving tasks for {ticket_key}") record_approval("task") return "task_router" @@ -122,7 +128,118 @@ def route_task_approval(state: WorkflowState) -> str: ) return END + # Handle standard (non-YOLO) approval draft ticket provisioning + # Note on split ownership: The gate nodes handle YOLO/autonomous paths where no manual comment or webhook + # is processed (so we must provision here), while the worker handles manual/human comment webhook triggers. + if not state.get("task_keys"): + from forge.integrations.jira.client import JiraClient + + jira = JiraClient() + try: + task_keys, tasks_by_repo = await provision_tasks_from_draft(state, jira) + # Store the newly created keys + state["task_keys"] = task_keys + state["tasks_by_repo"] = tasks_by_repo + except Exception as e: + logger.error( + f"Failed ticket provisioning during task approval for {ticket_key}: {e}", + exc_info=True, + ) + raise + finally: + await jira.close() + # Tasks approved, proceed to implementation logger.info(f"Tasks approved for {ticket_key}, proceeding to implementation") record_approval("task") return "task_router" + + +async def provision_tasks_from_draft( + state: WorkflowState, jira: "JiraClient" +) -> tuple[list[str], dict[str, list[str]]]: + """Provision Tasks from the task draft attachment on Jira. + + Args: + state: The workflow state dictionary. + jira: An active JiraClient instance. + + Returns: + Tuple of (task_keys, tasks_by_repo). + """ + ticket_key = state["ticket_key"] + from forge.config import get_settings + from forge.integrations.jira.client import MissingProjectConfig + from forge.models.workflow import ForgeLabel + from forge.workflow.utils.draft_manager import FORGE_TASKS_DRAFT_FILENAME, DraftManager + + settings = get_settings() + logger.info(f"Downloading task draft for {ticket_key}") + draft = await DraftManager.get_draft_attachment(jira, ticket_key, FORGE_TASKS_DRAFT_FILENAME) + if not draft: + raise ValueError(f"Approved draft {FORGE_TASKS_DRAFT_FILENAME} not found on {ticket_key}") + + parent_issue = await jira.get_issue(ticket_key) + project_key = parent_issue.project_key + + task_keys: list[str] = [] + tasks_by_repo: dict[str, list[str]] = {} + for item in draft.items: + if item.excluded: + logger.info(f"Skipping excluded task item {item.id}: {item.summary}") + continue + + # Fallback repository logic (mimics task_generation.py) + repo = item.repo + if not repo or repo == "unknown" or "/" not in repo: + try: + repo = await jira.get_project_default_repo(project_key) + except MissingProjectConfig: + repo = ( + settings.github_default_repo + if not settings.forge_require_project_config + else "" + ) + + if not repo or "/" not in repo: + logger.warning( + f"Task '{item.summary}' has no valid repo. " + "Set repo labels on Feature/Epic or GITHUB_DEFAULT_REPO." + ) + repo = "unknown" + + # Epic parent key logic: + # If draft item has epic_key set, use it. + # Else fallback to state's epic_keys. + epic_key = item.epic_key + if not epic_key and state.get("epic_keys"): + epic_key = state["epic_keys"][0] + + # Labels + labels = [ + ForgeLabel.FORGE_MANAGED.value, + f"forge:parent:{ticket_key}", + ] + if repo and repo != "unknown": + labels.append(f"repo:{repo}") + + task_key = await jira.create_task( + project_key=project_key, + summary=item.summary, + description=item.description, + parent_key=epic_key, + labels=labels, + ) + task_keys.append(task_key) + + if repo and repo != "unknown": + if repo not in tasks_by_repo: + tasks_by_repo[repo] = [] + tasks_by_repo[repo].append(task_key) + + # Delete the draft only after 100% successful ticket creation + await DraftManager.delete_draft_attachment(jira, ticket_key, FORGE_TASKS_DRAFT_FILENAME) + logger.info( + f"Successfully provisioned {len(task_keys)} Tasks across {len(tasks_by_repo)} repos and deleted draft" + ) + return task_keys, tasks_by_repo diff --git a/src/forge/workflow/nodes/epic_decomposition.py b/src/forge/workflow/nodes/epic_decomposition.py index ecafd9b4..d100bbfe 100644 --- a/src/forge/workflow/nodes/epic_decomposition.py +++ b/src/forge/workflow/nodes/epic_decomposition.py @@ -1,14 +1,17 @@ """Epic decomposition node for LangGraph workflow.""" import logging -from typing import Any +from datetime import UTC, datetime +from typing import Any, cast from forge.config import get_settings from forge.integrations.agents import ForgeAgent from forge.integrations.jira.client import JiraClient, MissingProjectConfig +from forge.models.draft import DraftItem, ForgeDecompositionDraft from forge.models.workflow import ForgeLabel from forge.workflow.feature.state import FeatureState as WorkflowState -from forge.workflow.utils import update_state_timestamp +from forge.workflow.utils import check_yolo_mode, update_state_timestamp +from forge.workflow.utils.draft_manager import FORGE_STORIES_DRAFT_FILENAME, DraftManager from forge.workflow.utils.jira_status import post_status_comment from forge.workflow.utils.qa_summary import post_qa_summary_if_needed @@ -85,18 +88,18 @@ async def decompose_epics(state: WorkflowState) -> WorkflowState: # 2. forge.repos Jira project property (required) feature_labels = await jira.get_labels(ticket_key) - available_repos = set() + available_repos_set: set[str] = set() # Add repos from Feature labels for label in feature_labels: if label.startswith("repo:"): - available_repos.add(label[5:]) + available_repos_set.add(label[5:]) # Add repos from Jira project property (required in strict mode) settings = get_settings() try: for repo in await jira.get_project_repos(project_key): - available_repos.add(repo) + available_repos_set.add(repo) except MissingProjectConfig as e: if settings.forge_require_project_config: logger.error( @@ -109,9 +112,9 @@ async def decompose_epics(state: WorkflowState) -> WorkflowState: return {**state, "last_error": str(e), "current_node": "decompose_epics"} logger.warning(f"Project {project_key}: {e} — falling back to GITHUB_KNOWN_REPOS") for repo in settings.known_repos: - available_repos.add(repo) + available_repos_set.add(repo) - available_repos = list(available_repos) + available_repos: list[str] = list(available_repos_set) # Build context for Epic generation context: dict[str, Any] = { @@ -132,73 +135,190 @@ async def decompose_epics(state: WorkflowState) -> WorkflowState: if not epics_data: logger.warning(f"No Epics generated for {ticket_key}") - return { - **state, - "last_error": "Epic generation returned no results", - "current_node": "decompose_epics", - } + return cast( + WorkflowState, + { + **state, + "last_error": "Epic generation returned no results", + "current_node": "decompose_epics", + "retry_count": state.get("retry_count", 0) + 1, + }, + ) - # Create Epics in Jira - secondary operation - epics_by_repo: dict[str, list[str]] = {} - - for epic in epics_data: - summary = epic.get("summary", "Untitled Epic") - plan = epic.get("plan", "") - repo = epic.get("repo", "") - - # Build labels for the Epic - # Include forge:managed for webhook routing and forge:parent for lookup - labels = [ - ForgeLabel.FORGE_MANAGED.value, - f"forge:parent:{ticket_key}", - ] - if repo and "/" in repo: - labels.append(f"repo:{repo}") - # Track which epics go to which repo - if repo not in epics_by_repo: - epics_by_repo[repo] = [] + # Check parent Jira ticket labels to check for forge:yolo and inspect global config yolo_mode + is_yolo = check_yolo_mode(state, feature_labels) - try: - epic_key = await jira.create_epic( - project_key=project_key, - summary=summary, - description=plan, - parent_key=ticket_key, - labels=labels, + if is_yolo: + # Create Epics in Jira immediately + epics_by_repo: dict[str, list[str]] = {} + + for epic in epics_data: + summary = epic.get("summary", "Untitled Epic") + plan = epic.get("plan", "") + repo = epic.get("repo", "") + + # Build labels for the Epic + # Include forge:managed for webhook routing and forge:parent for lookup + labels = [ + ForgeLabel.FORGE_MANAGED.value, + f"forge:parent:{ticket_key}", + ] + if repo and "/" in repo: + labels.append(f"repo:{repo}") + # Track which epics go to which repo + if repo not in epics_by_repo: + epics_by_repo[repo] = [] + + try: + epic_key = await jira.create_epic( + project_key=project_key, + summary=summary, + description=plan, + parent_key=ticket_key, + labels=labels, + ) + epic_keys.append(epic_key) + + if repo: + epics_by_repo[repo].append(epic_key) + + logger.info( + f"Created Epic {epic_key}: {summary}" + (f" (repo: {repo})" if repo else "") + ) + except Exception as e: + # Log but continue creating remaining Epics + jira_error = str(e) + logger.warning(f"Failed to create Epic '{summary}' for {ticket_key}: {e}") + + logger.info(f"Created {len(epic_keys)} Epics for {ticket_key}") + + # If we created some Epics, advance even with partial failures + if epic_keys: + # Only set workflow label after confirming epics were created + try: + await jira.set_workflow_label(ticket_key, ForgeLabel.PLAN_PENDING) + except Exception as e: + jira_error = str(e) + logger.warning(f"Failed to set workflow label for {ticket_key}: {e}") + + await jira.add_comment( + ticket_key, + "## 🤖 Forge interaction options\n\n" + f"- ✅ **Approve:** add `{ForgeLabel.PLAN_APPROVED.value}` to continue.\n" + "- ♻️ **Revise all epics:** add a comment starting with `!` on this ticket.\n" + "- 🔧 **Revise a single epic:** add a comment starting with `!` on the Epic.\n" + "- ❓ **Ask a question:** add a Jira comment starting with `?`.", ) - epic_keys.append(epic_key) - if repo: - epics_by_repo[repo].append(epic_key) + # Store plan summary in generation_context so Q&A can reference it + generation_context = state.get("generation_context", {}) + plan_summary_parts = [] + for epic in epics_data: + summary = epic.get("summary", "") + plan = epic.get("plan", "") + repo = epic.get("repo", "") + plan_summary_parts.append( + f"## {summary}" + (f" (repo: {repo})" if repo else "") + f"\n{plan}" + ) + generation_context["plan"] = "\n\n".join(plan_summary_parts) + + return cast( + WorkflowState, + update_state_timestamp( + { + **state, + "epic_keys": epic_keys, + "generation_context": generation_context, + "feedback_comment": None, + "revision_requested": False, + "current_epic_key": None, + "current_node": "plan_approval_gate", + "last_error": f"Partial Jira failure: {jira_error}" + if jira_error + else None, + } + ), + ) + else: + # No Epics created at all - this is a failure + return cast( + WorkflowState, + { + **state, + "last_error": jira_error or "Failed to create any Epics in Jira", + "current_node": "decompose_epics", + "retry_count": state.get("retry_count", 0) + 1, + }, + ) + else: + # Draft Review Flow (YOLO is inactive) + # Empty-draft guard to prevent proceeding without draft epics + if not epics_data: + return cast( + WorkflowState, + { + **state, + "last_error": "Failed to generate any draft Epics", + "current_node": "decompose_epics", + "retry_count": state.get("retry_count", 0) + 1, + }, + ) - logger.info( - f"Created Epic {epic_key}: {summary}" + (f" (repo: {repo})" if repo else "") + # Prior to saving, check for existing forge-stories-draft.json attachments + # and delete them using DraftManager/JiraClient to prevent duplicate file accumulation. + try: + await DraftManager.delete_draft_attachment( + jira, ticket_key, FORGE_STORIES_DRAFT_FILENAME ) except Exception as e: - # Log but continue creating remaining Epics - jira_error = str(e) - logger.warning(f"Failed to create Epic '{summary}' for {ticket_key}: {e}") + logger.warning(f"Failed to delete existing draft attachment: {e}") - logger.info(f"Created {len(epic_keys)} Epics for {ticket_key}") + # Convert epics_data into DraftItem instances + draft_items = [] + for idx, epic in enumerate(epics_data, start=1): + summary = epic.get("summary", "Untitled Epic") + plan = epic.get("plan", "") + repo = epic.get("repo", "") + draft_items.append( + DraftItem( + id=idx, + summary=summary, + description=plan, + repo=repo, + acceptance_criteria=[], + excluded=False, + ) + ) - # If we created some Epics, advance even with partial failures - if epic_keys: - # Only set workflow label after confirming epics were created + # Create Draft model + draft = ForgeDecompositionDraft( + parent_key=ticket_key, + phase="stories", + items=draft_items, + version=1, + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + + # Call DraftManager to serialize and save the generated epics draft to forge-stories-draft.json + await DraftManager.save_draft_attachment( + jira, ticket_key, draft, FORGE_STORIES_DRAFT_FILENAME + ) + + # Format Markdown review comment outlining proposed items + # Implement BR-003 Truncation Boundary + comment_body = DraftManager.format_review_comment(draft) + + # Post the review comment to the parent Jira ticket + await jira.add_comment(ticket_key, comment_body) + + # Set workflow label to pending try: await jira.set_workflow_label(ticket_key, ForgeLabel.PLAN_PENDING) except Exception as e: jira_error = str(e) logger.warning(f"Failed to set workflow label for {ticket_key}: {e}") - await jira.add_comment( - ticket_key, - "## 🤖 Forge interaction options\n\n" - f"- ✅ **Approve:** add `{ForgeLabel.PLAN_APPROVED.value}` to continue.\n" - "- ♻️ **Revise all epics:** add a comment starting with `!` on this ticket.\n" - "- 🔧 **Revise a single epic:** add a comment starting with `!` on the Epic.\n" - "- ❓ **Ask a question:** add a Jira comment starting with `?`.", - ) - # Store plan summary in generation_context so Q&A can reference it generation_context = state.get("generation_context", {}) plan_summary_parts = [] @@ -211,26 +331,23 @@ async def decompose_epics(state: WorkflowState) -> WorkflowState: ) generation_context["plan"] = "\n\n".join(plan_summary_parts) - return update_state_timestamp( - { - **state, - "epic_keys": epic_keys, - "generation_context": generation_context, - "feedback_comment": None, - "revision_requested": False, - "current_epic_key": None, - "current_node": "plan_approval_gate", - "last_error": f"Partial Jira failure: {jira_error}" if jira_error else None, - } + # Transition state to pause the workflow at the plan_approval_gate (setting is_paused = True and appropriate workflow flags) + return cast( + WorkflowState, + update_state_timestamp( + { + **state, + "epic_keys": [], + "generation_context": generation_context, + "feedback_comment": None, + "revision_requested": False, + "current_epic_key": None, + "current_node": "plan_approval_gate", + "is_paused": True, + "last_error": f"Partial Jira failure: {jira_error}" if jira_error else None, + } + ), ) - else: - # No Epics created at all - this is a failure - return { - **state, - "last_error": jira_error or "Failed to create any Epics in Jira", - "current_node": "decompose_epics", - "retry_count": state.get("retry_count", 0) + 1, - } except Exception as e: logger.error(f"Epic decomposition failed for {ticket_key}: {e}") @@ -243,7 +360,7 @@ async def decompose_epics(state: WorkflowState) -> WorkflowState: } if epic_keys: result_state["epic_keys"] = epic_keys - return result_state + return cast(WorkflowState, result_state) finally: await jira.close() await agent.close() @@ -286,16 +403,19 @@ async def regenerate_all_epics(state: WorkflowState) -> WorkflowState: } # Re-run decomposition (which will use context including feedback) - return await decompose_epics(updated_state) + return await decompose_epics(cast(WorkflowState, updated_state)) except Exception as e: logger.error(f"Epic regeneration failed for {ticket_key}: {e}") - return { - **state, - "last_error": str(e), - "current_node": "regenerate_all_epics", - "retry_count": state.get("retry_count", 0) + 1, - } + return cast( + WorkflowState, + { + **state, + "last_error": str(e), + "current_node": "regenerate_all_epics", + "retry_count": state.get("retry_count", 0) + 1, + }, + ) finally: await jira.close() @@ -313,7 +433,7 @@ async def update_single_epic(state: WorkflowState) -> WorkflowState: """ ticket_key = state["ticket_key"] epic_key = state.get("current_epic_key") - feedback = state.get("feedback_comment", "") + feedback = state.get("feedback_comment") or "" if not epic_key: logger.warning(f"No current_epic_key for single Epic update on {ticket_key}") @@ -356,25 +476,31 @@ async def update_single_epic(state: WorkflowState) -> WorkflowState: logger.info(f"Updated Epic {epic_key} plan") - return update_state_timestamp( - { - **state, - "current_epic_key": None, - "feedback_comment": None, - "revision_requested": False, - "current_node": "plan_approval_gate", - "last_error": None, - } + return cast( + WorkflowState, + update_state_timestamp( + { + **state, + "current_epic_key": None, + "feedback_comment": None, + "revision_requested": False, + "current_node": "plan_approval_gate", + "last_error": None, + } + ), ) except Exception as e: logger.error(f"Epic update failed for {epic_key}: {e}") - return { - **state, - "last_error": str(e), - "current_node": "update_single_epic", - "retry_count": state.get("retry_count", 0) + 1, - } + return cast( + WorkflowState, + { + **state, + "last_error": str(e), + "current_node": "update_single_epic", + "retry_count": state.get("retry_count", 0) + 1, + }, + ) finally: await jira.close() await agent.close() diff --git a/src/forge/workflow/nodes/task_generation.py b/src/forge/workflow/nodes/task_generation.py index 7c0d128c..9c1e9d6a 100644 --- a/src/forge/workflow/nodes/task_generation.py +++ b/src/forge/workflow/nodes/task_generation.py @@ -3,15 +3,18 @@ import asyncio import logging import re -from typing import Any +from datetime import UTC, datetime +from typing import Any, cast from forge.config import get_settings from forge.integrations.agents import ForgeAgent from forge.integrations.jira.client import JiraClient, MissingProjectConfig +from forge.models.draft import DraftItem, ForgeDecompositionDraft from forge.models.workflow import ForgeLabel from forge.prompts import load_prompt from forge.workflow.feature.state import FeatureState as WorkflowState -from forge.workflow.utils import update_state_timestamp +from forge.workflow.utils import check_yolo_mode, update_state_timestamp +from forge.workflow.utils.draft_manager import FORGE_TASKS_DRAFT_FILENAME, DraftManager from forge.workflow.utils.jira_status import post_status_comment logger = logging.getLogger(__name__) @@ -69,6 +72,9 @@ async def generate_tasks(state: WorkflowState) -> WorkflowState: # Get project key from parent Feature parent_issue = await jira.get_issue(ticket_key) project_key = parent_issue.project_key + feature_labels = await jira.get_labels(ticket_key) + + is_yolo = check_yolo_mode(state, feature_labels) # Pre-fetch all epic details upfront for sibling context for ek in epic_keys: @@ -85,6 +91,8 @@ async def generate_tasks(state: WorkflowState) -> WorkflowState: logger.warning(f"Failed to pre-fetch Epic {ek}: {e}") all_epics_details.append({"epic_key": ek, "epic_summary": ek, "epic_plan": ""}) + proposed_tasks_list = [] + for epic_key in epic_keys: logger.info(f"Generating Tasks for Epic {epic_key}") @@ -135,7 +143,7 @@ async def generate_tasks(state: WorkflowState) -> WorkflowState: existing_tasks=created_tasks_context if created_tasks_context else None, ) - # Create Tasks in Jira - secondary operation + # Create Tasks in Jira (YOLO) or collect (non-YOLO) for task in tasks_data: summary = task.get("summary", "Untitled Task") description = task.get("description", "") @@ -162,88 +170,206 @@ async def generate_tasks(state: WorkflowState) -> WorkflowState: ) repo = "unknown" - # Add labels: forge:managed for webhook routing, forge:parent for lookup, repo - labels = [ - ForgeLabel.FORGE_MANAGED.value, - f"forge:parent:{ticket_key}", # Parent Feature key - ] - if repo and repo != "unknown": - labels.append(f"repo:{repo}") - - try: - task_key = await jira.create_task( - project_key=project_key, - summary=summary, - description=description, - parent_key=epic_key, - labels=labels, - ) + if is_yolo: + # Add labels: forge:managed for webhook routing, forge:parent for lookup, repo + labels = [ + ForgeLabel.FORGE_MANAGED.value, + f"forge:parent:{ticket_key}", # Parent Feature key + ] + if repo and repo != "unknown": + labels.append(f"repo:{repo}") - all_task_keys.append(task_key) + try: + task_key = await jira.create_task( + project_key=project_key, + summary=summary, + description=description, + parent_key=epic_key, + labels=labels, + ) - # Track by repository - if repo not in tasks_by_repo: - tasks_by_repo[repo] = [] - tasks_by_repo[repo].append(task_key) + all_task_keys.append(task_key) + + # Track by repository + if repo not in tasks_by_repo: + tasks_by_repo[repo] = [] + tasks_by_repo[repo].append(task_key) + + # Track for context in subsequent epic task generation + created_tasks_context.append( + { + "epic_key": epic_key, + "epic_summary": epic_summary, + "task_key": task_key, + "summary": summary, + } + ) - # Track for context in subsequent epic task generation + logger.info(f"Created Task {task_key}: {summary} (repo: {repo})") + except Exception as e: + # Log but continue creating remaining Tasks + jira_error = str(e) + logger.warning(f"Failed to create Task '{summary}' for {ticket_key}: {e}") + else: + # Non-YOLO mode: collect proposed task details for draft + proposed_tasks_list.append( + { + "summary": summary, + "description": description, + "repo": repo, + "epic_key": epic_key, + } + ) + # Track for context in sibling generations + virtual_key = f"Draft Task {len(proposed_tasks_list)}" created_tasks_context.append( { "epic_key": epic_key, "epic_summary": epic_summary, - "task_key": task_key, + "task_key": virtual_key, "summary": summary, } ) - logger.info(f"Created Task {task_key}: {summary} (repo: {repo})") + if is_yolo: + logger.info( + f"Created {len(all_task_keys)} Tasks for {ticket_key}, awaiting implementation approval" + ) + + # If we created some Tasks, advance even with partial failures + if all_task_keys: + # Only set workflow label after confirming tasks were created + try: + await jira.set_workflow_label(ticket_key, ForgeLabel.TASK_PENDING) except Exception as e: - # Log but continue creating remaining Tasks jira_error = str(e) - logger.warning(f"Failed to create Task '{summary}' for {ticket_key}: {e}") + logger.warning(f"Failed to set workflow label for {ticket_key}: {e}") + + await jira.add_comment( + ticket_key, + "## 🤖 Forge interaction options\n\n" + f"- ✅ **Approve:** add `{ForgeLabel.TASK_APPROVED.value}` to continue.\n" + "- ♻️ **Revise all tasks:** add a comment starting with `!` on this ticket.\n" + "- 🔧 **Revise a single task:** add a comment starting with `!` on the Task.\n" + "- ❓ **Ask a question:** add a Jira comment starting with `?`.", + ) + return cast( + WorkflowState, + update_state_timestamp( + { + **state, + "task_keys": all_task_keys, + "tasks_by_repo": tasks_by_repo, + "feedback_comment": None, + "revision_requested": False, + "current_task_key": None, + "current_epic_key": None, + "current_node": "task_approval_gate", + "last_error": f"Partial Jira failure: {jira_error}" + if jira_error + else None, + } + ), + ) + else: + # No Tasks created at all - this is a failure + return cast( + WorkflowState, + { + **state, + "last_error": jira_error or "Failed to create any Tasks in Jira", + "current_node": "generate_tasks", + "retry_count": state.get("retry_count", 0) + 1, + }, + ) + else: + # Non-YOLO mode: Draft Review Flow + if not proposed_tasks_list: + return cast( + WorkflowState, + { + **state, + "last_error": "Failed to generate any draft Tasks", + "current_node": "generate_tasks", + "retry_count": state.get("retry_count", 0) + 1, + }, + ) - logger.info( - f"Created {len(all_task_keys)} Tasks for {ticket_key}, awaiting implementation approval" - ) + # Prior to saving, check for existing forge-tasks-draft.json attachments + # and delete them using DraftManager/JiraClient to prevent duplicate file accumulation. + try: + await DraftManager.delete_draft_attachment( + jira, ticket_key, FORGE_TASKS_DRAFT_FILENAME + ) + except Exception as e: + logger.warning(f"Failed to delete existing draft attachment: {e}") + + # Convert proposed_tasks_list into DraftItem instances + draft_items = [] + for idx, task_item in enumerate(proposed_tasks_list, start=1): + summary = task_item.get("summary", "Untitled Task") + description = task_item.get("description", "") + repo = task_item.get("repo", "unknown") + item_epic_key = task_item.get("epic_key") + draft_items.append( + DraftItem( + id=idx, + summary=summary, + description=description, + repo=repo, + epic_key=item_epic_key, + acceptance_criteria=[], + excluded=False, + ) + ) - # If we created some Tasks, advance even with partial failures - if all_task_keys: - # Only set workflow label after confirming tasks were created + # Create Draft model + draft = ForgeDecompositionDraft( + parent_key=ticket_key, + phase="tasks", + items=draft_items, + version=1, + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + + # Call DraftManager to serialize and save the generated tasks draft to forge-tasks-draft.json + await DraftManager.save_draft_attachment( + jira, ticket_key, draft, FORGE_TASKS_DRAFT_FILENAME + ) + + # Format Markdown review comment outlining proposed tasks + # Implement BR-003 Truncation Boundary + comment_body = DraftManager.format_review_comment(draft) + + # Post the review comment to the parent Jira ticket + await jira.add_comment(ticket_key, comment_body) + + # Set workflow label to pending try: await jira.set_workflow_label(ticket_key, ForgeLabel.TASK_PENDING) except Exception as e: jira_error = str(e) logger.warning(f"Failed to set workflow label for {ticket_key}: {e}") - await jira.add_comment( - ticket_key, - "## 🤖 Forge interaction options\n\n" - f"- ✅ **Approve:** add `{ForgeLabel.TASK_APPROVED.value}` to continue.\n" - "- ♻️ **Revise all tasks:** add a comment starting with `!` on this ticket.\n" - "- 🔧 **Revise a single task:** add a comment starting with `!` on the Task.\n" - "- ❓ **Ask a question:** add a Jira comment starting with `?`.", - ) - return update_state_timestamp( - { - **state, - "task_keys": all_task_keys, - "tasks_by_repo": tasks_by_repo, - "feedback_comment": None, - "revision_requested": False, - "current_task_key": None, - "current_epic_key": None, - "current_node": "task_approval_gate", - "last_error": f"Partial Jira failure: {jira_error}" if jira_error else None, - } + # Transition state to pause the workflow at the task_approval_gate (setting is_paused = True and appropriate workflow flags) + return cast( + WorkflowState, + update_state_timestamp( + { + **state, + "task_keys": [], + "tasks_by_repo": {}, + "feedback_comment": None, + "revision_requested": False, + "current_task_key": None, + "current_epic_key": None, + "current_node": "task_approval_gate", + "is_paused": True, + "last_error": f"Partial Jira failure: {jira_error}" if jira_error else None, + } + ), ) - else: - # No Tasks created at all - this is a failure - return { - **state, - "last_error": jira_error or "Failed to create any Tasks in Jira", - "current_node": "generate_tasks", - "retry_count": state.get("retry_count", 0) + 1, - } except Exception as e: logger.error(f"Task generation failed for {ticket_key}: {e}") @@ -257,9 +383,10 @@ async def generate_tasks(state: WorkflowState) -> WorkflowState: if all_task_keys: result_state["task_keys"] = all_task_keys result_state["tasks_by_repo"] = tasks_by_repo - return result_state + return cast(WorkflowState, result_state) finally: await jira.close() + await agent.close() async def _generate_tasks_for_epic( @@ -493,16 +620,19 @@ async def regenerate_all_tasks(state: WorkflowState) -> WorkflowState: } # Re-run task generation (which will incorporate feedback in context) - return await generate_tasks(updated_state) + return await generate_tasks(cast(WorkflowState, updated_state)) except Exception as e: logger.error(f"Task regeneration failed for {ticket_key}: {e}") - return { - **state, - "last_error": str(e), - "current_node": "regenerate_all_tasks", - "retry_count": state.get("retry_count", 0) + 1, - } + return cast( + WorkflowState, + { + **state, + "last_error": str(e), + "current_node": "regenerate_all_tasks", + "retry_count": state.get("retry_count", 0) + 1, + }, + ) finally: await jira.close() @@ -749,31 +879,37 @@ async def _fetch_sibling(ek: str) -> dict[str, str] | None: all_task_keys = remaining_task_keys + new_task_keys logger.info(f"Regenerated {len(new_task_keys)} tasks for Epic {epic_key} on {ticket_key}") - return update_state_timestamp( + return cast( + WorkflowState, + update_state_timestamp( + { + **state, + "task_keys": all_task_keys, + "tasks_by_repo": remaining_tasks_by_repo, + "feedback_comment": None, + "revision_requested": False, + "current_epic_key": None, + "current_node": "task_approval_gate", + "last_error": f"Partial Jira failure: {jira_error}" if jira_error else None, + } + ), + ) + + except Exception as e: + logger.error(f"Epic task regeneration failed for {epic_key} on {ticket_key}: {e}") + return cast( + WorkflowState, { **state, - "task_keys": all_task_keys, - "tasks_by_repo": remaining_tasks_by_repo, - "feedback_comment": None, + "last_error": str(e), + "current_node": "regenerate_epic_tasks", + "retry_count": state.get("retry_count", 0) + 1, + # Clear revision flags so task_approval_gate returns END instead of looping "revision_requested": False, + "feedback_comment": None, "current_epic_key": None, - "current_node": "task_approval_gate", - "last_error": f"Partial Jira failure: {jira_error}" if jira_error else None, - } + }, ) - - except Exception as e: - logger.error(f"Epic task regeneration failed for {epic_key} on {ticket_key}: {e}") - return { - **state, - "last_error": str(e), - "current_node": "regenerate_epic_tasks", - "retry_count": state.get("retry_count", 0) + 1, - # Clear revision flags so task_approval_gate returns END instead of looping - "revision_requested": False, - "feedback_comment": None, - "current_epic_key": None, - } finally: await jira.close() await agent.close() @@ -792,7 +928,7 @@ async def update_single_task(state: WorkflowState) -> WorkflowState: """ ticket_key = state["ticket_key"] task_key = state.get("current_task_key") - feedback = state.get("feedback_comment", "") + feedback = state.get("feedback_comment") or "" if not task_key: logger.warning(f"No current_task_key for single Task update on {ticket_key}") @@ -835,25 +971,31 @@ async def update_single_task(state: WorkflowState) -> WorkflowState: logger.info(f"Task {task_key} updated with feedback") - return update_state_timestamp( - { - **state, - "current_task_key": None, - "feedback_comment": None, - "revision_requested": False, - "current_node": "task_approval_gate", - "last_error": None, - } + return cast( + WorkflowState, + update_state_timestamp( + { + **state, + "current_task_key": None, + "feedback_comment": None, + "revision_requested": False, + "current_node": "task_approval_gate", + "last_error": None, + } + ), ) except Exception as e: logger.error(f"Task update failed for {task_key}: {e}") - return { - **state, - "last_error": str(e), - "current_node": "update_single_task", - "retry_count": state.get("retry_count", 0) + 1, - } + return cast( + WorkflowState, + { + **state, + "last_error": str(e), + "current_node": "update_single_task", + "retry_count": state.get("retry_count", 0) + 1, + }, + ) finally: await jira.close() await agent.close() diff --git a/src/forge/workflow/utils/__init__.py b/src/forge/workflow/utils/__init__.py index 9fadeb79..6360a5c3 100644 --- a/src/forge/workflow/utils/__init__.py +++ b/src/forge/workflow/utils/__init__.py @@ -5,7 +5,16 @@ from langgraph.graph import END -from forge.workflow.utils.comment_classifier import CommentType, classify_comment +from forge.workflow.utils.comment_classifier import ( + CommentType, + classify_comment, + parse_comment_command, +) +from forge.workflow.utils.draft_manager import ( + FORGE_STORIES_DRAFT_FILENAME, + FORGE_TASKS_DRAFT_FILENAME, + DraftManager, +) from forge.workflow.utils.jira_status import ( post_status_comment, remove_implementing_label, @@ -85,9 +94,34 @@ def set_error(state: dict[str, Any], error: str) -> dict[str, Any]: } +def check_yolo_mode(state: Any, labels: list[str] | None = None) -> bool: + """Check if YOLO mode is enabled based on labels, global settings, or state. + + The three components are: + 1. 'forge:yolo' label in the provided labels or the state context labels. + 2. Global configuration yolo_mode (settings.yolo_mode). + 3. State yolo_mode. + """ + from forge.config import get_settings + + settings = get_settings() + + # 1. Label check + has_label = False + if (labels and "forge:yolo" in labels) or "forge:yolo" in state.get("context", {}).get( + "labels", [] + ): + has_label = True + + # 2. Settings check & 3. State check + return has_label or settings.yolo_mode or bool(state.get("yolo_mode", False)) + + __all__ = [ + "check_yolo_mode", "CommentType", "classify_comment", + "parse_comment_command", "collect_review_exhaustion", "merge_review_exhaustion", "post_qa_summary_if_needed", @@ -102,4 +136,7 @@ def set_error(state: dict[str, Any], error: str) -> dict[str, Any]: "set_review_pending_label", "transition_tasks_to_in_progress", "update_state_timestamp", + "FORGE_STORIES_DRAFT_FILENAME", + "FORGE_TASKS_DRAFT_FILENAME", + "DraftManager", ] diff --git a/src/forge/workflow/utils/comment_classifier.py b/src/forge/workflow/utils/comment_classifier.py index 8caf9b5e..b583e973 100644 --- a/src/forge/workflow/utils/comment_classifier.py +++ b/src/forge/workflow/utils/comment_classifier.py @@ -2,6 +2,7 @@ import re from enum import StrEnum +from typing import Any class CommentType(StrEnum): @@ -10,6 +11,7 @@ class CommentType(StrEnum): QUESTION = "question" FEEDBACK = "feedback" INFORMATIONAL = "informational" + COMMAND = "command" # Legacy @forge ask pattern (case insensitive). @@ -21,12 +23,152 @@ class CommentType(StrEnum): # Pattern for revision prefix (allowing leading whitespace) _REVISION_PATTERN = re.compile(r"^\s*!") +# Regex to match case-insensitive /forge command prefix followed by command name +_FORGE_COMMAND_PATTERN = re.compile(r"^\s*/forge\s+([a-zA-Z0-9_-]+)", re.IGNORECASE) + +# Regex to match key-value pairs supporting single/double quoted string values or unquoted values +_SINGLE_PAIR_PATTERN = re.compile( + r'\s*([a-zA-Z_][a-zA-Z0-9_-]*)\s*=\s*(?:"([^"]*)"|\'([^\']*)\'|([^\s\'"]+))' +) + + +def _parse_key_values(args_text: str) -> dict[str, str]: + """Parse key-value pairs from argument text. + + Args: + args_text: The text to parse. + + Returns: + A dictionary of parsed parameters. + + Raises: + ValueError: If parameters are malformed. + """ + pos = 0 + params = {} + while pos < len(args_text): + m = _SINGLE_PAIR_PATTERN.match(args_text, pos) + if not m: + raise ValueError(f"Malformed parameters or trailing junk near: '{args_text[pos:]}'") + key = m.group(1) + val = ( + m.group(2) + if m.group(2) is not None + else (m.group(3) if m.group(3) is not None else m.group(4)) + ) + params[key] = val + pos = m.end() + return params + + +def parse_comment_command(comment_text: str) -> dict[str, Any] | None: + """Parse a /forge comment command and extract its parameters. + + Supported commands: + - remove: /forge remove + - add: /forge add key=val key2="val with spaces" + - update: /forge update key=val key2="val with spaces" + - exclude: /forge exclude + - approve: /forge approve + + Args: + comment_text: The comment text to parse. + + Returns: + A dictionary containing the parsed 'command' and arguments, + or an 'error' description if parameters are malformed, + or None if not a recognized /forge command. + """ + if not comment_text or not comment_text.strip(): + return None + + match = _FORGE_COMMAND_PATTERN.match(comment_text) + if not match: + return None + + cmd_name = match.group(1).lower() + valid_commands = {"remove", "add", "update", "exclude", "approve"} + if cmd_name not in valid_commands: + return None + + args_text = comment_text[match.end() :].strip() + + if cmd_name == "approve": + if args_text: + return { + "command": "approve", + "error": "approve command does not accept parameters", + } + return {"command": "approve"} + + if cmd_name in ("remove", "exclude"): + if not args_text: + return { + "command": cmd_name, + "error": f"Missing integer ID for {cmd_name} command", + } + if re.match(r"^\d+$", args_text): + return {"command": cmd_name, "id": int(args_text)} + return { + "command": cmd_name, + "error": f"Invalid integer ID for {cmd_name} command: '{args_text}'", + } + + if cmd_name == "add": + if not args_text: + return { + "command": "add", + "error": "Missing key-value parameters for add command", + } + try: + params = _parse_key_values(args_text) + except ValueError as e: + return { + "command": "add", + "error": str(e), + } + return {"command": "add", "params": params} + + if cmd_name == "update": + if not args_text: + return { + "command": "update", + "error": "Missing integer ID and parameters for update command", + } + id_match = re.match(r"^(\d+)(?:\s+(.*))?$", args_text) + if not id_match: + first_word = args_text.split(None, 1)[0] + if not re.match(r"^\d+$", first_word): + return { + "command": "update", + "error": f"Invalid integer ID for update command: '{first_word}'", + } + return { + "command": "update", + "error": "Missing integer ID for update command", + } + id_val = int(id_match.group(1)) + params_text = (id_match.group(2) or "").strip() + params = {} + if params_text: + try: + params = _parse_key_values(params_text) + except ValueError as e: + return { + "command": "update", + "error": str(e), + } + return {"command": "update", "id": id_val, "params": params} + + return None + def classify_comment(comment_text: str) -> CommentType: - """Classify a comment into question, feedback, or informational. + """Classify a comment into question, feedback, command, or informational. Classification rules: - - Questions: Comments starting with '?' + - Commands: Comments starting with /forge (except skip-gate/unskip-gate) + - Questions: Comments starting with '?' or '@forge ask' - Feedback (revision request): Comments starting with '!' - Informational: Everything else — ignored by the workflow @@ -42,6 +184,11 @@ def classify_comment(comment_text: str) -> CommentType: if not comment_text or not comment_text.strip(): return CommentType.INFORMATIONAL + # Check for commands first, since they are specific prefix patterns. + # Note: skip-gate/unskip-gate are excluded and should not return CommentType.COMMAND. + if parse_comment_command(comment_text) is not None: + return CommentType.COMMAND + if _QUESTION_MARK_PATTERN.match(comment_text): return CommentType.QUESTION diff --git a/src/forge/workflow/utils/draft_manager.py b/src/forge/workflow/utils/draft_manager.py new file mode 100644 index 00000000..3aa3f096 --- /dev/null +++ b/src/forge/workflow/utils/draft_manager.py @@ -0,0 +1,387 @@ +"""Utility for managing draft CRUD operations on Jira parent tickets as attachments.""" + +import copy +import logging +from typing import Any + +from pydantic import ValidationError + +from forge.integrations.jira import JiraClient +from forge.models.draft import ForgeDecompositionDraft + +logger = logging.getLogger(__name__) + +FORGE_STORIES_DRAFT_FILENAME = "forge-stories-draft.json" +FORGE_TASKS_DRAFT_FILENAME = "forge-tasks-draft.json" + + +class DraftManager: + """Manages draft CRUD operations on Jira parent tickets as attachments.""" + + @staticmethod + def _validate_item_params( + params: dict[str, Any], target_item: dict[str, Any] | None = None + ) -> None: + """Validate the fields in draft item parameters strictly. + + Args: + params: The parameters dictionary. + target_item: Optional target item dictionary to merge with (for update command). + + Raises: + ValueError: If a validation check fails. + """ + from forge.models.draft import DraftItem + + if target_item is not None: + full_item = {**target_item, **params} + else: + defaults = { + "id": 1, + "summary": "", + "description": "", + "repo": "", + "acceptance_criteria": [], + "excluded": False, + "epic_key": None, + } + full_item = {**defaults, **params} + + try: + DraftItem.model_validate(full_item, strict=True) + except ValidationError as e: + for error in e.errors(): + loc = error["loc"] + if not loc: + continue + field = str(loc[0]) + error_type = error["type"] + if error_type == "extra_forbidden": + raise ValueError(f"Unknown field '{field}'") + elif field in {"summary", "description", "repo"}: + val = ( + params.get(field) + if field in params + else (target_item.get(field) if target_item else None) + ) + raise ValueError( + f"Field '{field}' must be a string, got {type(val).__name__ if val is not None else 'None'}." + ) + elif field == "acceptance_criteria": + raise ValueError("Field 'acceptance_criteria' must be a list of strings.") + elif field == "excluded": + raise ValueError("Field 'excluded' must be a boolean.") + elif field == "epic_key": + raise ValueError("Field 'epic_key' must be a string or None.") + raise ValueError(str(e)) + + @staticmethod + def apply_draft_modification( + draft_json: list[dict[str, Any]], + parsed_command: dict[str, Any], + ) -> list[dict[str, Any]]: + """Apply a direct mutation on a list of draft story or task JSON objects based on the command type. + + Args: + draft_json: The current list of draft item dictionaries. + parsed_command: The parsed comment command dictionary. + + Returns: + The mutated list of draft item dictionaries. + + Raises: + ValueError: If the command contains an error, the target ID is missing/not found, + or strict type validation fails. + """ + if "error" in parsed_command: + raise ValueError(f"Invalid command parameters: {parsed_command['error']}") + + command = parsed_command.get("command") + if not command: + raise ValueError("Command type is missing in parsed command.") + + mutated_list = copy.deepcopy(draft_json) + + if command == "remove": + target_id = parsed_command.get("id") + if target_id is None: + raise ValueError("Missing ID for removal.") + + # Find and remove item + found = False + for i, item in enumerate(mutated_list): + if item.get("id") == target_id: + mutated_list.pop(i) + found = True + break + + if not found: + raise ValueError(f"Item with ID {target_id} not found for removal.") + + # Re-sequence remaining items + for idx, item in enumerate(mutated_list): + item["id"] = idx + 1 + + elif command == "add": + next_id = len(mutated_list) + 1 + params = parsed_command.get("params", {}) + + # Strict type validation + DraftManager._validate_item_params(params) + + # Build the new item using parsed parameters with defaults + new_item = { + "id": next_id, + "summary": params.get("summary", ""), + "description": params.get("description", ""), + "repo": params.get("repo", ""), + "acceptance_criteria": params.get("acceptance_criteria", []), + "excluded": params.get("excluded", False), + } + + mutated_list.append(new_item) + + elif command == "update": + target_id = parsed_command.get("id") + if target_id is None: + raise ValueError("Missing ID for update.") + + # Find the item + target_item = None + for item in mutated_list: + if item.get("id") == target_id: + target_item = item + break + + if not target_item: + raise ValueError(f"Item with ID {target_id} not found for update.") + + params = parsed_command.get("params", {}) + + # Strict type validation + DraftManager._validate_item_params(params, target_item) + + # Apply updates + for k, v in params.items(): + target_item[k] = v + + elif command == "exclude": + target_id = parsed_command.get("id") + if target_id is None: + raise ValueError("Missing ID for exclude command.") + + # Find the item + target_item = None + for item in mutated_list: + if item.get("id") == target_id: + target_item = item + break + + if not target_item: + raise ValueError(f"Item with ID {target_id} not found for exclude.") + + # Flip the excluded boolean key + target_item["excluded"] = not target_item.get("excluded", False) + + else: + raise ValueError(f"Unsupported modification command type: '{command}'") + + return mutated_list + + @staticmethod + async def save_draft_attachment( + jira_client: JiraClient, + issue_key: str, + draft: ForgeDecompositionDraft, + filename: str, + ) -> None: + """Save a draft decomposition as an attachment on a Jira parent issue, enforcing the single-file constraint. + + Args: + jira_client: The Jira client instance. + issue_key: The Jira issue key. + draft: The draft model to save. + filename: The target filename. + """ + # 1. Delete any matching filename to enforce the single-file constraint (BR-002/BR-004) + try: + await jira_client.delete_attachments_by_name(issue_key, filename) + except Exception as e: + logger.error( + f"Failed to delete existing draft attachment '{filename}' on {issue_key} to enforce single-file constraint: {e}", + exc_info=True, + ) + raise + + # 2. Serialize and upload + try: + content_json = draft.model_dump_json() + content_bytes = content_json.encode("utf-8") + except Exception as e: + logger.error(f"Failed to serialize draft for {issue_key}: {e}", exc_info=True) + raise + + try: + logger.info(f"Uploading new draft attachment '{filename}' to {issue_key}.") + await jira_client.add_attachment(issue_key, filename, content_bytes) + except Exception as e: + logger.error( + f"Failed to upload draft attachment '{filename}' to {issue_key}: {e}", + exc_info=True, + ) + raise + + @staticmethod + async def get_draft_attachment( + jira_client: JiraClient, + issue_key: str, + filename: str, + ) -> ForgeDecompositionDraft | None: + """Scan for attachment with matching filename on the issue, download and parse it. + + Args: + jira_client: The Jira client instance. + issue_key: The Jira issue key. + filename: The target filename to retrieve. + + Returns: + The parsed ForgeDecompositionDraft model instance, or None if not found or validation/parsing fails. + """ + try: + attachments = await jira_client.get_attachments(issue_key) + except Exception as e: + logger.error(f"Failed to list attachments for {issue_key}: {e}", exc_info=True) + raise + + target_attachment = None + for att in attachments: + if att.get("filename") == filename: + target_attachment = att + break + + if not target_attachment: + logger.debug(f"No attachment found with filename '{filename}' on {issue_key}") + return None + + content_url = target_attachment.get("content_url") or target_attachment.get("content") + if not content_url: + logger.warning( + f"Attachment '{filename}' found on {issue_key} but is missing a download URL." + ) + return None + + try: + content_bytes = await jira_client.download_attachment(content_url) + except Exception as e: + logger.error( + f"Failed to download attachment '{filename}' from {content_url}: {e}", + exc_info=True, + ) + raise + + try: + return ForgeDecompositionDraft.model_validate_json(content_bytes) + except ValidationError as ve: + if "json_invalid" in str(ve) or "Invalid JSON" in str(ve): + logger.warning( + f"Failed to parse draft attachment '{filename}' on {issue_key}. Error: {ve}", + exc_info=True, + ) + else: + logger.warning( + f"Validation failed for draft attachment '{filename}' on {issue_key}. Error: {ve}", + exc_info=True, + ) + return None + except Exception as e: + logger.warning( + f"Failed to parse draft attachment '{filename}' on {issue_key}. Error: {e}", + exc_info=True, + ) + return None + + @staticmethod + async def delete_draft_attachment( + jira_client: JiraClient, + issue_key: str, + filename: str, + ) -> None: + """Scan for any attachment with matching filename and delete it. + + Args: + jira_client: The Jira client instance. + issue_key: The Jira issue key. + filename: The target filename to delete. + """ + try: + await jira_client.delete_attachments_by_name(issue_key, filename) + except Exception as e: + logger.error( + f"Failed to delete draft attachments named '{filename}' on {issue_key}: {e}", + exc_info=True, + ) + raise + + @staticmethod + def format_review_comment(draft: ForgeDecompositionDraft) -> str: + """Format a human-readable review comment for a draft.""" + from forge.models.workflow import ForgeLabel + + items = draft.items + if draft.phase == "stories": + phase_title = "Epics" + noun_plural = "epics" + noun_singular = "epic" + phase_action = "decomposition" + item_label = "Plan" + approval_label = ForgeLabel.PLAN_APPROVED.value + filename = FORGE_STORIES_DRAFT_FILENAME + else: + phase_title = "Tasks" + noun_plural = "tasks" + noun_singular = "task" + phase_action = "implementation" + item_label = "Description" + approval_label = ForgeLabel.TASK_APPROVED.value + filename = FORGE_TASKS_DRAFT_FILENAME + + header = f"### 📋 Proposed {phase_title} Draft\n\nThe following {phase_title} have been proposed for {phase_action}:\n\n" + table = "| ID | Summary | Target Repo |\n|----|---------|-------------|\n" + for item in items: + table += f"| {item.id} | {item.summary} | {item.repo or 'unknown'} |\n" + table += "\n---\n\n" + + details = "" + for item in items: + details += f"#### {item.id}. {item.summary} (Repo: {item.repo or 'unknown'})\n" + if item.description: + details += f"**{item_label}:**\n{item.description}\n\n" + else: + details += "\n" + + footer = ( + "## 🤖 Forge interaction options\n\n" + f"- ✅ **Approve:** comment `/forge approve` or add `{approval_label}` to continue.\n" + f"- ♻️ **Revise all {noun_plural}:** add a comment starting with `!` on this ticket.\n" + f"- 🔧 **Revise a single {noun_singular}:** add a comment starting with `!` on the {phase_title.rstrip('s')}.\n" + "- ❓ **Ask a question:** add a Jira comment starting with `?`." + ) + + full_comment = header + table + details + footer + + if len(full_comment) > 32767 or len(items) > 15: + condensed_table = "| ID | Summary | Target Repo |\n|----|---------|-------------|\n" + for item in items: + condensed_table += f"| {item.id} | {item.summary} | {item.repo or 'unknown'} |\n" + + condensed_comment = ( + f"### 📋 Proposed {phase_title} Draft (Condensed)\n\n" + "⚠️ **Warning:** The proposed plan exceeds character or size limits for detailed display in a comment. " + f"Please refer to the attached `{filename}` for full implementation plan details.\n\n" + + condensed_table + + "\n" + + footer + ) + return condensed_comment + + return full_comment diff --git a/tests/contracts/test_github_contracts.py b/tests/contracts/test_github_contracts.py index d42bbb6c..a26223fd 100644 --- a/tests/contracts/test_github_contracts.py +++ b/tests/contracts/test_github_contracts.py @@ -114,10 +114,10 @@ def test_parse_pull_request_merged(self): "merged": True, "title": "PROJ-104: OAuth implementation", "head": {"ref": "feature/PROJ-104"}, - "html_url": "https://github.com/acme/backend/pull/42" + "html_url": "https://github.com/acme/backend/pull/42", }, "repository": {"full_name": "acme/backend"}, - "sender": {"login": "senior-dev"} + "sender": {"login": "senior-dev"}, } data = parse_github_webhook( payload=payload, @@ -139,10 +139,10 @@ def test_parse_pull_request_closed_not_merged(self): "merged": False, "title": "WIP: Experimental feature", "head": {"ref": "feature/experiment"}, - "html_url": "https://github.com/acme/backend/pull/43" + "html_url": "https://github.com/acme/backend/pull/43", }, "repository": {"full_name": "acme/backend"}, - "sender": {"login": "dev-user"} + "sender": {"login": "dev-user"}, } data = parse_github_webhook( payload=payload, @@ -186,17 +186,17 @@ def test_parse_pr_review_changes_requested(self): "user": {"login": "senior-dev"}, "body": "Please add error handling for the token refresh.", "state": "changes_requested", - "submitted_at": "2024-03-20T15:00:00Z" + "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" + "html_url": "https://github.com/acme/backend/pull/42", }, "repository": {"full_name": "acme/backend"}, - "sender": {"login": "senior-dev"} + "sender": {"login": "senior-dev"}, } data = parse_github_webhook( payload=payload, @@ -220,10 +220,10 @@ def test_extract_from_pr_title(self): "state": "open", "title": "[PROJ-123] Fix login bug", "head": {"ref": "fix-login"}, - "html_url": "https://github.com/org/repo/pull/1" + "html_url": "https://github.com/org/repo/pull/1", }, "repository": {"full_name": "org/repo"}, - "sender": {"login": "user"} + "sender": {"login": "user"}, } data = parse_github_webhook(payload, "pull_request", "id-1") assert data.ticket_key == "PROJ-123" @@ -237,10 +237,10 @@ def test_extract_from_branch_when_title_has_no_ticket(self): "state": "open", "title": "Fix login bug", "head": {"ref": "feature/PROJ-456-login"}, - "html_url": "https://github.com/org/repo/pull/1" + "html_url": "https://github.com/org/repo/pull/1", }, "repository": {"full_name": "org/repo"}, - "sender": {"login": "user"} + "sender": {"login": "user"}, } data = parse_github_webhook(payload, "pull_request", "id-2") assert data.ticket_key == "PROJ-456" @@ -266,10 +266,10 @@ def test_extract_ticket_various_formats(self): "state": "open", "title": text, "head": {"ref": "main"}, - "html_url": "https://github.com/org/repo/pull/1" + "html_url": "https://github.com/org/repo/pull/1", }, "repository": {"full_name": "org/repo"}, - "sender": {"login": "user"} + "sender": {"login": "user"}, } data = parse_github_webhook(payload, "pull_request", "id") assert data.ticket_key == expected_key, f"Failed for: {text}" @@ -283,10 +283,10 @@ def test_no_ticket_found(self): "state": "open", "title": "Fix some bug", "head": {"ref": "fix-bug"}, - "html_url": "https://github.com/org/repo/pull/1" + "html_url": "https://github.com/org/repo/pull/1", }, "repository": {"full_name": "org/repo"}, - "sender": {"login": "user"} + "sender": {"login": "user"}, } data = parse_github_webhook(payload, "pull_request", "id") assert data.ticket_key is None @@ -302,7 +302,7 @@ def test_parse_push_with_ticket_in_branch(self): "after": "newcommitsha123456789012345678901234", "before": "oldcommitsha123456789012345678901234", "repository": {"full_name": "acme/backend"}, - "sender": {"login": "developer"} + "sender": {"login": "developer"}, } data = parse_github_webhook(payload, "push", "delivery-push-1") @@ -317,11 +317,7 @@ class TestEdgeCases: def test_minimal_payload(self): """Handle minimal payload with missing optional fields.""" - payload = { - "action": "created", - "repository": {}, - "sender": {} - } + payload = {"action": "created", "repository": {}, "sender": {}} data = parse_github_webhook(payload, "unknown", "id-1") assert data.event_type == "unknown" @@ -340,10 +336,10 @@ def test_check_run_without_pull_requests(self): "status": "completed", "conclusion": "success", "head_sha": "sha123", - "pull_requests": [] # No associated PRs + "pull_requests": [], # No associated PRs }, "repository": {"full_name": "acme/repo"}, - "sender": {"login": "bot"} + "sender": {"login": "bot"}, } data = parse_github_webhook(payload, "check_run", "id-1") @@ -358,7 +354,7 @@ def test_raw_payload_preserved(self): "action": "opened", "custom_field": "custom_value", "repository": {"full_name": "org/repo"}, - "sender": {"login": "user"} + "sender": {"login": "user"}, } data = parse_github_webhook(payload, "test", "id-1") @@ -377,16 +373,11 @@ def test_parse_pr_comment(self): "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..." + "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"} + "sender": {"login": "reviewer"}, } data = parse_github_webhook(payload, "issue_comment", "id-1") @@ -403,15 +394,12 @@ def test_parse_issue_comment_not_pr(self): "issue": { "number": 100, "title": "Bug report", - "html_url": "https://github.com/acme/backend/issues/100" + "html_url": "https://github.com/acme/backend/issues/100", # No pull_request field }, - "comment": { - "id": 12346, - "body": "Can you provide more details?" - }, + "comment": {"id": 12346, "body": "Can you provide more details?"}, "repository": {"full_name": "acme/backend"}, - "sender": {"login": "maintainer"} + "sender": {"login": "maintainer"}, } data = parse_github_webhook(payload, "issue_comment", "id-2") diff --git a/tests/contracts/test_jira_contracts.py b/tests/contracts/test_jira_contracts.py index b84f66ae..1b5229a6 100644 --- a/tests/contracts/test_jira_contracts.py +++ b/tests/contracts/test_jira_contracts.py @@ -110,8 +110,8 @@ def test_parse_issue_with_missing_optional_fields(self): "fields": { "issuetype": {"name": "Task"}, "status": {"name": "Open"}, - "summary": "Minimal issue" - } + "summary": "Minimal issue", + }, } issue = JiraIssue.from_api_response(minimal_data) @@ -134,12 +134,8 @@ def test_parse_issue_with_empty_adf_content(self): "issuetype": {"name": "Feature"}, "status": {"name": "New"}, "summary": "Test", - "description": { - "version": 1, - "type": "doc", - "content": [] - } - } + "description": {"version": 1, "type": "doc", "content": []}, + }, } issue = JiraIssue.from_api_response(data) @@ -170,13 +166,10 @@ def test_parse_comment_with_plain_text_body(self): """Parse a comment with plain text body.""" data = { "id": "10200", - "author": { - "accountId": "user-123", - "displayName": "Bob Smith" - }, + "author": {"accountId": "user-123", "displayName": "Bob Smith"}, "body": "LGTM! Approved.", "created": "2024-03-21T10:00:00.000+0000", - "updated": "2024-03-21T10:00:00.000+0000" + "updated": "2024-03-21T10:00:00.000+0000", } comment = JiraComment.from_api_response(data) @@ -186,11 +179,7 @@ def test_parse_comment_with_plain_text_body(self): def test_parse_comment_with_missing_optional_fields(self): """Parse a comment with minimal fields.""" - data = { - "id": "10300", - "author": {}, - "body": "Simple comment" - } + data = {"id": "10300", "author": {}, "body": "Simple comment"} comment = JiraComment.from_api_response(data) assert comment.id == "10300" @@ -209,19 +198,9 @@ def test_extract_text_from_nested_paragraphs(self): "version": 1, "type": "doc", "content": [ - { - "type": "paragraph", - "content": [ - {"type": "text", "text": "First paragraph."} - ] - }, - { - "type": "paragraph", - "content": [ - {"type": "text", "text": "Second paragraph."} - ] - } - ] + {"type": "paragraph", "content": [{"type": "text", "text": "First paragraph."}]}, + {"type": "paragraph", "content": [{"type": "text", "text": "Second paragraph."}]}, + ], } text = JiraIssue._extract_text_from_adf(adf) assert "First paragraph." in text @@ -247,15 +226,13 @@ def test_extract_text_from_bullet_list(self): "content": [ { "type": "paragraph", - "content": [ - {"type": "text", "text": "Item one"} - ] + "content": [{"type": "text", "text": "Item one"}], } - ] + ], } - ] + ], } - ] + ], } # Note: Current implementation may not handle nested list items # This test documents the current behavior @@ -270,12 +247,7 @@ class TestProjectKeyExtraction: def test_extract_project_key_standard(self): """Extract project key from standard issue key.""" issue = JiraIssue( - key="PROJ-123", - id="1", - summary="Test", - description="", - status="Open", - issue_type="Task" + key="PROJ-123", id="1", summary="Test", description="", status="Open", issue_type="Task" ) assert issue.project_key == "PROJ" @@ -287,7 +259,7 @@ def test_extract_project_key_multi_part(self): summary="Test", description="", status="Open", - issue_type="Task" + issue_type="Task", ) # Should extract everything before the last hyphen-number assert issue.project_key == "MY-PROJECT" @@ -295,12 +267,7 @@ def test_extract_project_key_multi_part(self): def test_extract_project_key_no_hyphen(self): """Handle key without hyphen (edge case).""" issue = JiraIssue( - key="INVALID", - id="1", - summary="Test", - description="", - status="Open", - issue_type="Task" + key="INVALID", id="1", summary="Test", description="", status="Open", issue_type="Task" ) assert issue.project_key == "INVALID" @@ -318,8 +285,8 @@ def test_parse_iso_date_with_timezone(self): "status": {"name": "Open"}, "summary": "Test", "created": "2024-03-15T10:23:45.000+0000", - "updated": "2024-03-20T14:30:22.000-0800" - } + "updated": "2024-03-20T14:30:22.000-0800", + }, } issue = JiraIssue.from_api_response(data) @@ -339,8 +306,8 @@ def test_parse_iso_date_with_z_suffix(self): "issuetype": {"name": "Task"}, "status": {"name": "Open"}, "summary": "Test", - "created": "2024-01-01T00:00:00.000Z" - } + "created": "2024-01-01T00:00:00.000Z", + }, } issue = JiraIssue.from_api_response(data) diff --git a/tests/flows/bug_workflow/test_complete_bug_flow.py b/tests/flows/bug_workflow/test_complete_bug_flow.py index 41c999c2..70ead3bb 100644 --- a/tests/flows/bug_workflow/test_complete_bug_flow.py +++ b/tests/flows/bug_workflow/test_complete_bug_flow.py @@ -134,22 +134,25 @@ def test_error_at_retry_cap_escalates(self): class TestBugWorkflowResumeRouting: """route_entry correctly resumes a bug workflow at any node.""" - @pytest.mark.parametrize("node,expected", [ - ("analyze_bug", "analyze_bug"), - ("regenerate_rca", "regenerate_rca"), # reruns cleanup+setup before analyze_bug - ("rca_approval_gate", "rca_option_gate"), # backward compat: old gate maps to new - ("setup_workspace", "setup_workspace"), - ("implement_bug_fix", "implement_bug_fix"), - ("create_pr", "create_pr"), - ("teardown_workspace", "teardown_workspace"), - ("ci_evaluator", "ci_evaluator"), - ("attempt_ci_fix", "ci_evaluator"), - ("wait_for_ci_gate", "wait_for_ci_gate"), - ("local_review", "local_review"), - ("ai_review", "human_review_gate"), - ("human_review_gate", "human_review_gate"), - ("escalate_blocked", "escalate_blocked"), - ]) + @pytest.mark.parametrize( + "node,expected", + [ + ("analyze_bug", "analyze_bug"), + ("regenerate_rca", "regenerate_rca"), # reruns cleanup+setup before analyze_bug + ("rca_approval_gate", "rca_option_gate"), # backward compat: old gate maps to new + ("setup_workspace", "setup_workspace"), + ("implement_bug_fix", "implement_bug_fix"), + ("create_pr", "create_pr"), + ("teardown_workspace", "teardown_workspace"), + ("ci_evaluator", "ci_evaluator"), + ("attempt_ci_fix", "ci_evaluator"), + ("wait_for_ci_gate", "wait_for_ci_gate"), + ("local_review", "local_review"), + ("ai_review", "human_review_gate"), + ("human_review_gate", "human_review_gate"), + ("escalate_blocked", "escalate_blocked"), + ], + ) def test_resume_routing(self, node, expected): """route_entry maps each node to the correct resume target.""" state = make_workflow_state( @@ -161,8 +164,7 @@ def test_resume_routing(self, node, expected): result = route_entry(state) assert result == expected, ( - f"route_entry with current_node='{node}' returned '{result}', " - f"expected '{expected}'" + f"route_entry with current_node='{node}' returned '{result}', expected '{expected}'" ) @@ -192,9 +194,15 @@ def test_minimal_old_state_without_new_fields_does_not_crash(self): def test_all_new_current_node_values_are_handled(self): """Every new current_node value from the redesign has a route_entry mapping.""" new_nodes = [ - "triage_check", "triage_gate", "reflect_rca", - "rca_option_gate", "plan_bug_fix", "plan_approval_gate", - "regenerate_plan", "decompose_plan", "post_merge_summary", + "triage_check", + "triage_gate", + "reflect_rca", + "rca_option_gate", + "plan_bug_fix", + "plan_approval_gate", + "regenerate_plan", + "decompose_plan", + "post_merge_summary", ] for node in new_nodes: state = make_workflow_state( @@ -230,18 +238,21 @@ def test_bug_plan_pending_routes_to_plan_approval_gate(self): class TestNewResumeRoutingCases: """New pipeline nodes resume correctly at the right point.""" - @pytest.mark.parametrize("node,expected", [ - ("triage_check", "triage_check"), - ("triage_gate", "triage_gate"), - ("reflect_rca", "reflect_rca"), - ("rca_option_gate", "rca_option_gate"), - ("plan_bug_fix", "plan_bug_fix"), - ("plan_approval_gate", "plan_approval_gate"), - ("regenerate_plan", "regenerate_plan"), - ("decompose_plan", "decompose_plan"), - ("post_merge_summary", "post_merge_summary"), - ("rca_approval_gate", "rca_option_gate"), # backward compat - ]) + @pytest.mark.parametrize( + "node,expected", + [ + ("triage_check", "triage_check"), + ("triage_gate", "triage_gate"), + ("reflect_rca", "reflect_rca"), + ("rca_option_gate", "rca_option_gate"), + ("plan_bug_fix", "plan_bug_fix"), + ("plan_approval_gate", "plan_approval_gate"), + ("regenerate_plan", "regenerate_plan"), + ("decompose_plan", "decompose_plan"), + ("post_merge_summary", "post_merge_summary"), + ("rca_approval_gate", "rca_option_gate"), # backward compat + ], + ) def test_resume_routing_new_pipeline_nodes(self, node, expected): """route_entry maps each new current_node to the correct resume target.""" state = make_workflow_state( @@ -276,11 +287,13 @@ async def test_missing_fields_pauses_at_triage_gate(self): mock_jira = MagicMock() mock_jira.add_comment = AsyncMock() mock_jira.set_workflow_label = AsyncMock() - mock_jira.get_issue = AsyncMock(return_value=MagicMock( - summary="Login fails", - description="Short desc", - project_key="BUG", - )) + mock_jira.get_issue = AsyncMock( + return_value=MagicMock( + summary="Login fails", + description="Short desc", + project_key="BUG", + ) + ) mock_jira.get_comments = AsyncMock(return_value=[]) mock_jira.close = AsyncMock() @@ -312,10 +325,13 @@ async def test_sufficient_ticket_routes_to_analyze_bug(self): mock_jira = MagicMock() mock_jira.add_comment = AsyncMock() - mock_jira.get_issue = AsyncMock(return_value=MagicMock( - summary="Login fails with $", description="Full description with all fields", - project_key="BUG", - )) + mock_jira.get_issue = AsyncMock( + return_value=MagicMock( + summary="Login fails with $", + description="Full description with all fields", + project_key="BUG", + ) + ) mock_jira.get_comments = AsyncMock(return_value=[]) mock_jira.close = AsyncMock() @@ -347,7 +363,9 @@ async def test_three_failed_reflections_routes_to_rca_option_gate(self): ticket_type=TicketType.BUG, is_paused=False, rca_content="## Root Cause\nBug is in validators.py", - rca_options=[{"title": "Fix regex", "description": "Update pattern", "tradeoffs": "Low risk"}], + rca_options=[ + {"title": "Fix regex", "description": "Update pattern", "tradeoffs": "Low risk"} + ], reflection_count=2, # Will become 3 after this run reflection_critique=None, ) @@ -383,6 +401,7 @@ class TestQualitativeRetryCapFlow: def test_qualitative_retry_count_two_routes_to_create_pr(self): """_route_after_local_review with qualitative_retry_count=2 → create_pr.""" from forge.workflow.bug.graph import _route_after_local_review + state = make_workflow_state( ticket_key="BUG-Q1", current_node="local_review", @@ -395,6 +414,7 @@ def test_qualitative_retry_count_two_routes_to_create_pr(self): def test_symptom_only_first_retry_routes_to_implement(self): """_route_after_local_review with symptom_only + retry=0 → implement_bug_fix.""" from forge.workflow.bug.graph import _route_after_local_review + state = make_workflow_state( ticket_key="BUG-Q2", current_node="local_review", @@ -415,25 +435,33 @@ class TestRouteAfterTriageCheck: def test_missing_fields_routes_to_triage_gate(self): state = make_workflow_state( - ticket_key="BUG-TC1", ticket_type=TicketType.BUG, current_node="triage_gate", + ticket_key="BUG-TC1", + ticket_type=TicketType.BUG, + current_node="triage_gate", ) assert _route_after_triage_check(state) == "triage_gate" def test_sufficient_ticket_routes_to_analyze_bug(self): state = make_workflow_state( - ticket_key="BUG-TC2", ticket_type=TicketType.BUG, current_node="analyze_bug", + ticket_key="BUG-TC2", + ticket_type=TicketType.BUG, + current_node="analyze_bug", ) assert _route_after_triage_check(state) == "analyze_bug" def test_error_routes_to_escalate_blocked(self): state = make_workflow_state( - ticket_key="BUG-TC3", ticket_type=TicketType.BUG, current_node="escalate_blocked", + ticket_key="BUG-TC3", + ticket_type=TicketType.BUG, + current_node="escalate_blocked", ) assert _route_after_triage_check(state) == "escalate_blocked" def test_unknown_node_defaults_to_triage_gate(self): state = make_workflow_state( - ticket_key="BUG-TC4", ticket_type=TicketType.BUG, current_node="something_unknown", + ticket_key="BUG-TC4", + ticket_type=TicketType.BUG, + current_node="something_unknown", ) assert _route_after_triage_check(state) == "triage_gate" @@ -443,19 +471,25 @@ class TestRouteAfterAnalyzeBug: def test_success_routes_to_reflect_rca(self): state = make_workflow_state( - ticket_key="BUG-AB1", ticket_type=TicketType.BUG, current_node="reflect_rca", + ticket_key="BUG-AB1", + ticket_type=TicketType.BUG, + current_node="reflect_rca", ) assert _route_after_analyze_bug(state) == "reflect_rca" def test_too_many_failures_routes_to_escalate(self): state = make_workflow_state( - ticket_key="BUG-AB2", ticket_type=TicketType.BUG, current_node="escalate_blocked", + ticket_key="BUG-AB2", + ticket_type=TicketType.BUG, + current_node="escalate_blocked", ) assert _route_after_analyze_bug(state) == "escalate_blocked" def test_container_failure_terminates_invocation(self): state = make_workflow_state( - ticket_key="BUG-AB3", ticket_type=TicketType.BUG, current_node="analyze_bug", + ticket_key="BUG-AB3", + ticket_type=TicketType.BUG, + current_node="analyze_bug", ) assert _route_after_analyze_bug(state) == END @@ -465,48 +499,67 @@ class TestRouteAfterReflectRca: def test_failure_state_routes_to_escalate(self): state = make_workflow_state( - ticket_key="BUG-RR1", ticket_type=TicketType.BUG, current_node="escalate_blocked", + ticket_key="BUG-RR1", + ticket_type=TicketType.BUG, + current_node="escalate_blocked", ) assert _route_after_reflect_rca(state) == "escalate_blocked" def test_container_failure_terminates(self): state = make_workflow_state( - ticket_key="BUG-RR2", ticket_type=TicketType.BUG, current_node="reflect_rca", + ticket_key="BUG-RR2", + ticket_type=TicketType.BUG, + current_node="reflect_rca", ) assert _route_after_reflect_rca(state) == END def test_reflection_cap_routes_to_rca_option_gate(self): state = make_workflow_state( - ticket_key="BUG-RR3", ticket_type=TicketType.BUG, current_node="rca_option_gate", - reflection_count=3, reflection_critique="still needs depth", + ticket_key="BUG-RR3", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", + reflection_count=3, + reflection_critique="still needs depth", ) assert _route_after_reflect_rca(state) == "rca_option_gate" def test_critique_below_cap_loops_to_analyze_bug(self): state = make_workflow_state( - ticket_key="BUG-RR4", ticket_type=TicketType.BUG, current_node="rca_option_gate", - reflection_count=1, reflection_critique="needs more depth on auth flow", + ticket_key="BUG-RR4", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", + reflection_count=1, + reflection_critique="needs more depth on auth flow", ) assert _route_after_reflect_rca(state) == "analyze_bug" def test_no_critique_routes_to_rca_option_gate(self): state = make_workflow_state( - ticket_key="BUG-RR5", ticket_type=TicketType.BUG, current_node="rca_option_gate", - reflection_count=1, reflection_critique=None, + ticket_key="BUG-RR5", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", + reflection_count=1, + reflection_critique=None, ) assert _route_after_reflect_rca(state) == "rca_option_gate" def test_empty_critique_routes_to_rca_option_gate(self): state = make_workflow_state( - ticket_key="BUG-RR6", ticket_type=TicketType.BUG, current_node="rca_option_gate", - reflection_count=1, reflection_critique="", + ticket_key="BUG-RR6", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", + reflection_count=1, + reflection_critique="", ) assert _route_after_reflect_rca(state) == "rca_option_gate" def test_whitespace_only_critique_routes_to_rca_option_gate(self): state = make_workflow_state( - ticket_key="BUG-RR7", ticket_type=TicketType.BUG, current_node="rca_option_gate", - reflection_count=1, reflection_critique=" ", + ticket_key="BUG-RR7", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", + reflection_count=1, + reflection_critique=" ", ) assert _route_after_reflect_rca(state) == "rca_option_gate" @@ -516,49 +569,68 @@ class TestRouteRcaOption: def test_question_routes_to_answer_question(self): state = make_workflow_state( - ticket_key="BUG-RO1", ticket_type=TicketType.BUG, current_node="rca_option_gate", + ticket_key="BUG-RO1", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", is_question=True, ) assert route_rca_option(state) == "answer_question" def test_question_takes_priority_over_selection(self): state = make_workflow_state( - ticket_key="BUG-RO2", ticket_type=TicketType.BUG, current_node="rca_option_gate", - is_question=True, selected_fix_option=1, is_paused=False, + ticket_key="BUG-RO2", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", + is_question=True, + selected_fix_option=1, + is_paused=False, ) assert route_rca_option(state) == "answer_question" def test_option_selected_routes_to_plan_bug_fix(self): state = make_workflow_state( - ticket_key="BUG-RO3", ticket_type=TicketType.BUG, current_node="rca_option_gate", - selected_fix_option=1, is_paused=False, + ticket_key="BUG-RO3", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", + selected_fix_option=1, + is_paused=False, ) assert route_rca_option(state) == "plan_bug_fix" def test_option_selected_while_paused_routes_to_end(self): state = make_workflow_state( - ticket_key="BUG-RO4", ticket_type=TicketType.BUG, current_node="rca_option_gate", - selected_fix_option=1, is_paused=True, + ticket_key="BUG-RO4", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", + selected_fix_option=1, + is_paused=True, ) assert route_rca_option(state) == END def test_revision_requested_routes_to_regenerate_rca(self): state = make_workflow_state( - ticket_key="BUG-RO5", ticket_type=TicketType.BUG, current_node="rca_option_gate", - revision_requested=True, is_paused=False, + ticket_key="BUG-RO5", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", + revision_requested=True, + is_paused=False, ) assert route_rca_option(state) == "regenerate_rca" def test_paused_routes_to_end(self): state = make_workflow_state( - ticket_key="BUG-RO6", ticket_type=TicketType.BUG, current_node="rca_option_gate", + ticket_key="BUG-RO6", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", is_paused=True, ) assert route_rca_option(state) == END def test_no_signals_routes_to_end(self): state = make_workflow_state( - ticket_key="BUG-RO7", ticket_type=TicketType.BUG, current_node="rca_option_gate", + ticket_key="BUG-RO7", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", is_paused=False, ) assert route_rca_option(state) == END @@ -569,36 +641,49 @@ class TestRoutePlanApproval: def test_question_routes_to_answer_question(self): state = make_workflow_state( - ticket_key="BUG-PA1", ticket_type=TicketType.BUG, current_node="plan_approval_gate", + ticket_key="BUG-PA1", + ticket_type=TicketType.BUG, + current_node="plan_approval_gate", is_question=True, ) assert route_plan_approval(state) == "answer_question" def test_paused_routes_to_end(self): state = make_workflow_state( - ticket_key="BUG-PA2", ticket_type=TicketType.BUG, current_node="plan_approval_gate", + ticket_key="BUG-PA2", + ticket_type=TicketType.BUG, + current_node="plan_approval_gate", is_paused=True, ) assert route_plan_approval(state) == END def test_revision_requested_routes_to_regenerate_plan(self): state = make_workflow_state( - ticket_key="BUG-PA3", ticket_type=TicketType.BUG, current_node="plan_approval_gate", - revision_requested=True, is_paused=False, + ticket_key="BUG-PA3", + ticket_type=TicketType.BUG, + current_node="plan_approval_gate", + revision_requested=True, + is_paused=False, ) assert route_plan_approval(state) == "regenerate_plan" def test_approved_routes_to_decompose_plan(self): state = make_workflow_state( - ticket_key="BUG-PA4", ticket_type=TicketType.BUG, current_node="plan_approval_gate", - is_paused=False, revision_requested=False, + ticket_key="BUG-PA4", + ticket_type=TicketType.BUG, + current_node="plan_approval_gate", + is_paused=False, + revision_requested=False, ) assert route_plan_approval(state) == "decompose_plan" def test_question_takes_priority_over_paused(self): state = make_workflow_state( - ticket_key="BUG-PA5", ticket_type=TicketType.BUG, current_node="plan_approval_gate", - is_question=True, is_paused=True, + ticket_key="BUG-PA5", + ticket_type=TicketType.BUG, + current_node="plan_approval_gate", + is_question=True, + is_paused=True, ) assert route_plan_approval(state) == "answer_question" @@ -608,29 +693,41 @@ class TestRouteAfterWorkspaceSetup: def test_success_routes_to_implement(self): state = make_workflow_state( - ticket_key="BUG-WS1", ticket_type=TicketType.BUG, current_node="setup_workspace", - workspace_path="/tmp/forge-ws", last_error=None, + ticket_key="BUG-WS1", + ticket_type=TicketType.BUG, + current_node="setup_workspace", + workspace_path="/tmp/forge-ws", + last_error=None, ) assert _route_after_workspace_setup(state) == "implement_bug_fix" def test_no_workspace_path_escalates(self): state = make_workflow_state( - ticket_key="BUG-WS2", ticket_type=TicketType.BUG, current_node="setup_workspace", - workspace_path=None, last_error=None, + ticket_key="BUG-WS2", + ticket_type=TicketType.BUG, + current_node="setup_workspace", + workspace_path=None, + last_error=None, ) assert _route_after_workspace_setup(state) == "escalate_blocked" def test_error_escalates(self): state = make_workflow_state( - ticket_key="BUG-WS3", ticket_type=TicketType.BUG, current_node="setup_workspace", - workspace_path="/tmp/forge-ws", last_error="clone failed", + ticket_key="BUG-WS3", + ticket_type=TicketType.BUG, + current_node="setup_workspace", + workspace_path="/tmp/forge-ws", + last_error="clone failed", ) assert _route_after_workspace_setup(state) == "escalate_blocked" def test_empty_workspace_path_escalates(self): state = make_workflow_state( - ticket_key="BUG-WS4", ticket_type=TicketType.BUG, current_node="setup_workspace", - workspace_path="", last_error=None, + ticket_key="BUG-WS4", + ticket_type=TicketType.BUG, + current_node="setup_workspace", + workspace_path="", + last_error=None, ) assert _route_after_workspace_setup(state) == "escalate_blocked" @@ -640,36 +737,51 @@ class TestRouteAfterImplementation: def test_no_error_routes_to_local_review(self): state = make_workflow_state( - ticket_key="BUG-IM1", ticket_type=TicketType.BUG, current_node="implement_bug_fix", - last_error=None, retry_count=0, + ticket_key="BUG-IM1", + ticket_type=TicketType.BUG, + current_node="implement_bug_fix", + last_error=None, + retry_count=0, ) assert _route_after_implementation(state) == "local_review" def test_error_below_cap_retries(self): state = make_workflow_state( - ticket_key="BUG-IM2", ticket_type=TicketType.BUG, current_node="implement_bug_fix", - last_error="timeout", retry_count=1, + ticket_key="BUG-IM2", + ticket_type=TicketType.BUG, + current_node="implement_bug_fix", + last_error="timeout", + retry_count=1, ) assert _route_after_implementation(state) == "implement_bug_fix" def test_error_at_cap_escalates(self): state = make_workflow_state( - ticket_key="BUG-IM3", ticket_type=TicketType.BUG, current_node="implement_bug_fix", - last_error="timeout", retry_count=3, + ticket_key="BUG-IM3", + ticket_type=TicketType.BUG, + current_node="implement_bug_fix", + last_error="timeout", + retry_count=3, ) assert _route_after_implementation(state) == "escalate_blocked" def test_error_above_cap_escalates(self): state = make_workflow_state( - ticket_key="BUG-IM4", ticket_type=TicketType.BUG, current_node="implement_bug_fix", - last_error="timeout", retry_count=5, + ticket_key="BUG-IM4", + ticket_type=TicketType.BUG, + current_node="implement_bug_fix", + last_error="timeout", + retry_count=5, ) assert _route_after_implementation(state) == "escalate_blocked" def test_no_error_ignores_high_retry_count(self): state = make_workflow_state( - ticket_key="BUG-IM5", ticket_type=TicketType.BUG, current_node="implement_bug_fix", - last_error=None, retry_count=5, + ticket_key="BUG-IM5", + ticket_type=TicketType.BUG, + current_node="implement_bug_fix", + last_error=None, + retry_count=5, ) assert _route_after_implementation(state) == "local_review" @@ -679,43 +791,61 @@ class TestRouteAfterLocalReview: def test_adequate_verdict_routes_to_update_docs(self): state = make_workflow_state( - ticket_key="BUG-LR1", ticket_type=TicketType.BUG, current_node="local_review", - local_review_verdict="adequate", qualitative_retry_count=0, + ticket_key="BUG-LR1", + ticket_type=TicketType.BUG, + current_node="local_review", + local_review_verdict="adequate", + qualitative_retry_count=0, ) assert _route_after_local_review(state) == "update_documentation" def test_tests_incomplete_routes_to_implement(self): state = make_workflow_state( - ticket_key="BUG-LR2", ticket_type=TicketType.BUG, current_node="local_review", - local_review_verdict="tests_incomplete", qualitative_retry_count=0, + ticket_key="BUG-LR2", + ticket_type=TicketType.BUG, + current_node="local_review", + local_review_verdict="tests_incomplete", + qualitative_retry_count=0, ) assert _route_after_local_review(state) == "implement_bug_fix" def test_symptom_only_routes_to_implement(self): state = make_workflow_state( - ticket_key="BUG-LR3", ticket_type=TicketType.BUG, current_node="local_review", - local_review_verdict="symptom_only", qualitative_retry_count=0, + ticket_key="BUG-LR3", + ticket_type=TicketType.BUG, + current_node="local_review", + local_review_verdict="symptom_only", + qualitative_retry_count=0, ) assert _route_after_local_review(state) == "implement_bug_fix" def test_tests_incomplete_at_cap_routes_to_update_docs(self): state = make_workflow_state( - ticket_key="BUG-LR4", ticket_type=TicketType.BUG, current_node="local_review", - local_review_verdict="tests_incomplete", qualitative_retry_count=2, + ticket_key="BUG-LR4", + ticket_type=TicketType.BUG, + current_node="local_review", + local_review_verdict="tests_incomplete", + qualitative_retry_count=2, ) assert _route_after_local_review(state) == "update_documentation" def test_no_verdict_mechanical_at_cap_routes_to_update_docs(self): state = make_workflow_state( - ticket_key="BUG-LR5", ticket_type=TicketType.BUG, current_node="local_review", - local_review_verdict=None, local_review_attempts=2, + ticket_key="BUG-LR5", + ticket_type=TicketType.BUG, + current_node="local_review", + local_review_verdict=None, + local_review_attempts=2, ) assert _route_after_local_review(state) == "update_documentation" def test_no_verdict_mechanical_below_cap_falls_back_to_current_node(self): state = make_workflow_state( - ticket_key="BUG-LR6", ticket_type=TicketType.BUG, current_node="local_review", - local_review_verdict=None, local_review_attempts=0, + ticket_key="BUG-LR6", + ticket_type=TicketType.BUG, + current_node="local_review", + local_review_verdict=None, + local_review_attempts=0, ) assert _route_after_local_review(state) == "local_review" @@ -725,29 +855,41 @@ class TestRouteAfterPrCreation: def test_success_routes_to_teardown(self): state = make_workflow_state( - ticket_key="BUG-PR1", ticket_type=TicketType.BUG, current_node="create_pr", - last_error=None, pr_urls=["https://github.com/org/repo/pull/1"], + ticket_key="BUG-PR1", + ticket_type=TicketType.BUG, + current_node="create_pr", + last_error=None, + pr_urls=["https://github.com/org/repo/pull/1"], ) assert _route_after_pr_creation(state) == "teardown_workspace" def test_error_with_no_pr_urls_escalates(self): state = make_workflow_state( - ticket_key="BUG-PR2", ticket_type=TicketType.BUG, current_node="create_pr", - last_error="PR creation failed", pr_urls=[], + ticket_key="BUG-PR2", + ticket_type=TicketType.BUG, + current_node="create_pr", + last_error="PR creation failed", + pr_urls=[], ) assert _route_after_pr_creation(state) == "escalate_blocked" def test_error_with_existing_pr_urls_routes_to_teardown(self): state = make_workflow_state( - ticket_key="BUG-PR3", ticket_type=TicketType.BUG, current_node="create_pr", - last_error="partial failure", pr_urls=["https://github.com/org/repo/pull/1"], + ticket_key="BUG-PR3", + ticket_type=TicketType.BUG, + current_node="create_pr", + last_error="partial failure", + pr_urls=["https://github.com/org/repo/pull/1"], ) assert _route_after_pr_creation(state) == "teardown_workspace" def test_no_error_no_pr_urls_routes_to_teardown(self): state = make_workflow_state( - ticket_key="BUG-PR4", ticket_type=TicketType.BUG, current_node="create_pr", - last_error=None, pr_urls=[], + ticket_key="BUG-PR4", + ticket_type=TicketType.BUG, + current_node="create_pr", + last_error=None, + pr_urls=[], ) assert _route_after_pr_creation(state) == "teardown_workspace" @@ -757,29 +899,41 @@ class TestRouteAfterTeardown: def test_remaining_repos_loops_to_setup_workspace(self): state = make_workflow_state( - ticket_key="BUG-TD1", ticket_type=TicketType.BUG, current_node="teardown_workspace", - repos_to_process=["org/a", "org/b"], repos_completed=["org/a"], + ticket_key="BUG-TD1", + ticket_type=TicketType.BUG, + current_node="teardown_workspace", + repos_to_process=["org/a", "org/b"], + repos_completed=["org/a"], ) assert _route_after_teardown(state) == "setup_workspace" def test_all_repos_done_routes_to_wait_for_ci_gate(self): state = make_workflow_state( - ticket_key="BUG-TD2", ticket_type=TicketType.BUG, current_node="teardown_workspace", - repos_to_process=["org/a"], repos_completed=["org/a"], + ticket_key="BUG-TD2", + ticket_type=TicketType.BUG, + current_node="teardown_workspace", + repos_to_process=["org/a"], + repos_completed=["org/a"], ) assert _route_after_teardown(state) == "wait_for_ci_gate" def test_empty_repos_routes_to_wait_for_ci_gate(self): state = make_workflow_state( - ticket_key="BUG-TD3", ticket_type=TicketType.BUG, current_node="teardown_workspace", - repos_to_process=[], repos_completed=[], + ticket_key="BUG-TD3", + ticket_type=TicketType.BUG, + current_node="teardown_workspace", + repos_to_process=[], + repos_completed=[], ) assert _route_after_teardown(state) == "wait_for_ci_gate" def test_multiple_remaining_repos_loops(self): state = make_workflow_state( - ticket_key="BUG-TD4", ticket_type=TicketType.BUG, current_node="teardown_workspace", - repos_to_process=["org/a", "org/b", "org/c"], repos_completed=[], + ticket_key="BUG-TD4", + ticket_type=TicketType.BUG, + current_node="teardown_workspace", + repos_to_process=["org/a", "org/b", "org/c"], + repos_completed=[], ) assert _route_after_teardown(state) == "setup_workspace" @@ -789,35 +943,45 @@ class TestRouteCiEvaluation: def test_passed_routes_to_human_review_gate(self): state = make_workflow_state( - ticket_key="BUG-CI1", ticket_type=TicketType.BUG, current_node="ci_evaluator", + ticket_key="BUG-CI1", + ticket_type=TicketType.BUG, + current_node="ci_evaluator", ci_status="passed", ) assert _route_ci_evaluation(state) == "human_review_gate" def test_fixing_routes_to_attempt_ci_fix(self): state = make_workflow_state( - ticket_key="BUG-CI2", ticket_type=TicketType.BUG, current_node="ci_evaluator", + ticket_key="BUG-CI2", + ticket_type=TicketType.BUG, + current_node="ci_evaluator", ci_status="fixing", ) assert _route_ci_evaluation(state) == "attempt_ci_fix" def test_pending_routes_to_end(self): state = make_workflow_state( - ticket_key="BUG-CI3", ticket_type=TicketType.BUG, current_node="ci_evaluator", + ticket_key="BUG-CI3", + ticket_type=TicketType.BUG, + current_node="ci_evaluator", ci_status="pending", ) assert _route_ci_evaluation(state) == END def test_failed_routes_to_escalate_blocked(self): state = make_workflow_state( - ticket_key="BUG-CI4", ticket_type=TicketType.BUG, current_node="ci_evaluator", + ticket_key="BUG-CI4", + ticket_type=TicketType.BUG, + current_node="ci_evaluator", ci_status="failed", ) assert _route_ci_evaluation(state) == "escalate_blocked" def test_empty_status_routes_to_escalate_blocked(self): state = make_workflow_state( - ticket_key="BUG-CI5", ticket_type=TicketType.BUG, current_node="ci_evaluator", + ticket_key="BUG-CI5", + ticket_type=TicketType.BUG, + current_node="ci_evaluator", ci_status="", ) assert _route_ci_evaluation(state) == "escalate_blocked" @@ -828,36 +992,53 @@ class TestRouteHumanReviewBug: def test_pr_merged_routes_to_post_merge_summary(self): state = make_workflow_state( - ticket_key="BUG-HR1", ticket_type=TicketType.BUG, current_node="human_review_gate", + ticket_key="BUG-HR1", + ticket_type=TicketType.BUG, + current_node="human_review_gate", pr_merged=True, ) assert _route_human_review_bug(state) == "post_merge_summary" def test_revision_requested_routes_to_implement_review(self): state = make_workflow_state( - ticket_key="BUG-HR2", ticket_type=TicketType.BUG, current_node="human_review_gate", - pr_merged=False, revision_requested=True, feedback_comment="fix the tests", + ticket_key="BUG-HR2", + ticket_type=TicketType.BUG, + current_node="human_review_gate", + pr_merged=False, + revision_requested=True, + feedback_comment="fix the tests", ) assert _route_human_review_bug(state) == "implement_review" def test_paused_routes_to_end(self): state = make_workflow_state( - ticket_key="BUG-HR3", ticket_type=TicketType.BUG, current_node="human_review_gate", - pr_merged=False, is_paused=True, + ticket_key="BUG-HR3", + ticket_type=TicketType.BUG, + current_node="human_review_gate", + pr_merged=False, + is_paused=True, ) assert _route_human_review_bug(state) == END def test_not_merged_not_paused_routes_to_complete_tasks(self): state = make_workflow_state( - ticket_key="BUG-HR4", ticket_type=TicketType.BUG, current_node="human_review_gate", - pr_merged=False, is_paused=False, revision_requested=False, + ticket_key="BUG-HR4", + ticket_type=TicketType.BUG, + current_node="human_review_gate", + pr_merged=False, + is_paused=False, + revision_requested=False, ) assert _route_human_review_bug(state) == "complete_tasks" def test_pr_merged_takes_priority_over_revision(self): state = make_workflow_state( - ticket_key="BUG-HR5", ticket_type=TicketType.BUG, current_node="human_review_gate", - pr_merged=True, revision_requested=True, feedback_comment="fix", + ticket_key="BUG-HR5", + ticket_type=TicketType.BUG, + current_node="human_review_gate", + pr_merged=True, + revision_requested=True, + feedback_comment="fix", ) assert _route_human_review_bug(state) == "post_merge_summary" @@ -867,30 +1048,40 @@ class TestRouteAfterAnswerBug: def test_returns_to_triage_gate(self): state = make_workflow_state( - ticket_key="BUG-AQ1", ticket_type=TicketType.BUG, current_node="triage_gate", + ticket_key="BUG-AQ1", + ticket_type=TicketType.BUG, + current_node="triage_gate", ) assert _route_after_answer_bug(state) == "triage_gate" def test_returns_to_rca_option_gate(self): state = make_workflow_state( - ticket_key="BUG-AQ2", ticket_type=TicketType.BUG, current_node="rca_option_gate", + ticket_key="BUG-AQ2", + ticket_type=TicketType.BUG, + current_node="rca_option_gate", ) assert _route_after_answer_bug(state) == "rca_option_gate" def test_returns_to_plan_approval_gate(self): state = make_workflow_state( - ticket_key="BUG-AQ3", ticket_type=TicketType.BUG, current_node="plan_approval_gate", + ticket_key="BUG-AQ3", + ticket_type=TicketType.BUG, + current_node="plan_approval_gate", ) assert _route_after_answer_bug(state) == "plan_approval_gate" def test_unknown_node_defaults_to_rca_option_gate(self): state = make_workflow_state( - ticket_key="BUG-AQ4", ticket_type=TicketType.BUG, current_node="implement_bug_fix", + ticket_key="BUG-AQ4", + ticket_type=TicketType.BUG, + current_node="implement_bug_fix", ) assert _route_after_answer_bug(state) == "rca_option_gate" def test_empty_node_defaults_to_rca_option_gate(self): state = make_workflow_state( - ticket_key="BUG-AQ5", ticket_type=TicketType.BUG, current_node="", + ticket_key="BUG-AQ5", + ticket_type=TicketType.BUG, + current_node="", ) assert _route_after_answer_bug(state) == "rca_option_gate" diff --git a/tests/flows/ci_recovery/test_ci_failure_and_fix.py b/tests/flows/ci_recovery/test_ci_failure_and_fix.py index 42515f17..646dba20 100644 --- a/tests/flows/ci_recovery/test_ci_failure_and_fix.py +++ b/tests/flows/ci_recovery/test_ci_failure_and_fix.py @@ -160,7 +160,7 @@ def test_ci_exhaustion_escalates_scenario(self): """ state = make_workflow_state( current_node="ci_evaluator", - ci_status="failed", # evaluator sets 'failed' after exhaustion + ci_status="failed", # evaluator sets 'failed' after exhaustion ci_fix_attempt=5, ci_failed_checks=[{"name": "lint", "conclusion": "failure"}], ) diff --git a/tests/flows/error_recovery/test_blocked_and_retry.py b/tests/flows/error_recovery/test_blocked_and_retry.py index 9521a014..3a576a84 100644 --- a/tests/flows/error_recovery/test_blocked_and_retry.py +++ b/tests/flows/error_recovery/test_blocked_and_retry.py @@ -1,6 +1,5 @@ """Flow tests for blocked state escalation and forge:retry recovery.""" - from forge.models.workflow import TicketType from forge.workflow.bug.graph import route_entry from forge.workflow.feature.graph import route_by_ticket_type @@ -73,9 +72,8 @@ def test_blocked_workflow_skips_invocation(self): state["is_blocked"] = True terminal_nodes = ("complete", "complete_tasks", "aggregate_feature_status") - is_terminal_or_blocked = ( - state.get("current_node") in terminal_nodes - or state.get("is_blocked", False) + is_terminal_or_blocked = state.get("current_node") in terminal_nodes or state.get( + "is_blocked", False ) assert is_terminal_or_blocked is True @@ -93,9 +91,8 @@ def test_mid_workflow_node_is_not_terminal(self): state["is_blocked"] = False terminal_nodes = ("complete", "complete_tasks", "aggregate_feature_status") - is_terminal_or_blocked = ( - state.get("current_node") in terminal_nodes - or state.get("is_blocked", False) + is_terminal_or_blocked = state.get("current_node") in terminal_nodes or state.get( + "is_blocked", False ) assert is_terminal_or_blocked is False diff --git a/tests/flows/feature_workflow/test_complete_feature_flow.py b/tests/flows/feature_workflow/test_complete_feature_flow.py index da8aafd1..2541eac3 100644 --- a/tests/flows/feature_workflow/test_complete_feature_flow.py +++ b/tests/flows/feature_workflow/test_complete_feature_flow.py @@ -1,6 +1,5 @@ """Tests for complete feature workflow flow.""" - import pytest from forge.models.workflow import TicketType @@ -66,6 +65,7 @@ def test_prd_approved_to_spec_generation(self): ) from forge.workflow.gates import route_prd_approval + next_node = route_prd_approval(state) assert next_node == "generate_spec" @@ -81,11 +81,13 @@ def test_spec_approved_to_epic_decomposition(self): ) from forge.workflow.gates import route_spec_approval + next_node = route_spec_approval(state) assert next_node == "decompose_epics" - def test_plan_approved_to_task_generation(self): + @pytest.mark.asyncio + async def test_plan_approved_to_task_generation(self): """Approved plan progresses to task generation when resumed.""" state = make_workflow_state( ticket_key="TEST-123", @@ -95,7 +97,8 @@ def test_plan_approved_to_task_generation(self): ) from forge.workflow.gates import route_plan_approval - next_node = route_plan_approval(state) + + next_node = await route_plan_approval(state) assert next_node == "generate_tasks" @@ -142,7 +145,8 @@ def test_multiple_epics_created(self, multi_epic_state): """Multiple epics are tracked in state.""" assert len(multi_epic_state["epic_keys"]) == 4 - def test_all_epics_must_be_approved(self, multi_epic_state): + @pytest.mark.asyncio + async def test_all_epics_must_be_approved(self, multi_epic_state): """All epics must be approved for plan approval - workflow pauses to wait.""" # Workflow is paused waiting for approval multi_epic_state["is_paused"] = True @@ -151,7 +155,7 @@ def test_all_epics_must_be_approved(self, multi_epic_state): from forge.workflow.gates import route_plan_approval - result = route_plan_approval(multi_epic_state) + result = await route_plan_approval(multi_epic_state) # Should wait (END) until approved via webhook assert result == END @@ -195,7 +199,8 @@ def test_all_repos_must_complete(self, multi_repo_state): # Should have more repos to process remaining = [ - r for r in multi_repo_state["repos_to_process"] + r + for r in multi_repo_state["repos_to_process"] if r not in multi_repo_state["repos_completed"] ] diff --git a/tests/flows/parallel_execution/test_task_routing.py b/tests/flows/parallel_execution/test_task_routing.py index 28db4778..1c26bce3 100644 --- a/tests/flows/parallel_execution/test_task_routing.py +++ b/tests/flows/parallel_execution/test_task_routing.py @@ -26,8 +26,11 @@ async def test_single_repo_initialises_state(self): tasks_by_repo={"org/backend": ["TEST-200", "TEST-201"]}, ) - with patch("forge.workflow.nodes.task_router.update_state_timestamp", side_effect=lambda s: s): + with patch( + "forge.workflow.nodes.task_router.update_state_timestamp", side_effect=lambda s: s + ): from forge.workflow.nodes.task_router import route_tasks_by_repo + result = await route_tasks_by_repo(state) assert result["repos_to_process"] == ["org/backend"] @@ -46,8 +49,11 @@ async def test_multi_repo_sets_first_repo_as_current(self): }, ) - with patch("forge.workflow.nodes.task_router.update_state_timestamp", side_effect=lambda s: s): + with patch( + "forge.workflow.nodes.task_router.update_state_timestamp", side_effect=lambda s: s + ): from forge.workflow.nodes.task_router import route_tasks_by_repo + result = await route_tasks_by_repo(state) assert len(result["repos_to_process"]) == 2 @@ -62,8 +68,11 @@ async def test_empty_tasks_by_repo_sets_error(self): tasks_by_repo={}, ) - with patch("forge.workflow.nodes.task_router.update_state_timestamp", side_effect=lambda s: s): + with patch( + "forge.workflow.nodes.task_router.update_state_timestamp", side_effect=lambda s: s + ): from forge.workflow.nodes.task_router import route_tasks_by_repo + result = await route_tasks_by_repo(state) assert result["last_error"] is not None diff --git a/tests/flows/status_transitions/test_label_transitions.py b/tests/flows/status_transitions/test_label_transitions.py index 1ae209ad..1ded49c7 100644 --- a/tests/flows/status_transitions/test_label_transitions.py +++ b/tests/flows/status_transitions/test_label_transitions.py @@ -1,6 +1,5 @@ """Tests for label state transitions.""" - import pytest from forge.models.workflow import ForgeLabel, get_workflow_phase @@ -163,28 +162,31 @@ def test_all_workflow_labels_start_with_forge(self): class TestLabelStateAtEachPhase: """Tests verifying correct label at each workflow phase.""" - @pytest.mark.parametrize("label,expected_phase", [ - (ForgeLabel.PRD_DRAFTING, "prd_generation"), - (ForgeLabel.PRD_PENDING, "prd_approval"), - (ForgeLabel.PRD_APPROVED, "spec_generation"), - (ForgeLabel.SPEC_DRAFTING, "spec_generation"), - (ForgeLabel.SPEC_PENDING, "spec_approval"), - (ForgeLabel.SPEC_APPROVED, "epic_decomposition"), - (ForgeLabel.PLAN_DRAFTING, "epic_decomposition"), - (ForgeLabel.PLAN_PENDING, "plan_approval"), - (ForgeLabel.PLAN_APPROVED, "task_generation"), - (ForgeLabel.TASK_GENERATED, "task_routing"), - (ForgeLabel.TASK_IMPLEMENTING, "implementation"), - (ForgeLabel.TASK_PR_CREATED, "pr_created"), - (ForgeLabel.TASK_CI_PENDING, "ci_evaluation"), - (ForgeLabel.TASK_CI_FAILED, "ci_fix"), - (ForgeLabel.TASK_REVIEW_PENDING, "human_review"), - (ForgeLabel.TASK_REVIEW_APPROVED, "complete"), - (ForgeLabel.RCA_DRAFTING, "rca_generation"), - (ForgeLabel.RCA_PENDING, "rca_approval"), - (ForgeLabel.RCA_APPROVED, "bug_fix"), - (ForgeLabel.BLOCKED, "blocked"), - ]) + @pytest.mark.parametrize( + "label,expected_phase", + [ + (ForgeLabel.PRD_DRAFTING, "prd_generation"), + (ForgeLabel.PRD_PENDING, "prd_approval"), + (ForgeLabel.PRD_APPROVED, "spec_generation"), + (ForgeLabel.SPEC_DRAFTING, "spec_generation"), + (ForgeLabel.SPEC_PENDING, "spec_approval"), + (ForgeLabel.SPEC_APPROVED, "epic_decomposition"), + (ForgeLabel.PLAN_DRAFTING, "epic_decomposition"), + (ForgeLabel.PLAN_PENDING, "plan_approval"), + (ForgeLabel.PLAN_APPROVED, "task_generation"), + (ForgeLabel.TASK_GENERATED, "task_routing"), + (ForgeLabel.TASK_IMPLEMENTING, "implementation"), + (ForgeLabel.TASK_PR_CREATED, "pr_created"), + (ForgeLabel.TASK_CI_PENDING, "ci_evaluation"), + (ForgeLabel.TASK_CI_FAILED, "ci_fix"), + (ForgeLabel.TASK_REVIEW_PENDING, "human_review"), + (ForgeLabel.TASK_REVIEW_APPROVED, "complete"), + (ForgeLabel.RCA_DRAFTING, "rca_generation"), + (ForgeLabel.RCA_PENDING, "rca_approval"), + (ForgeLabel.RCA_APPROVED, "bug_fix"), + (ForgeLabel.BLOCKED, "blocked"), + ], + ) def test_label_maps_to_phase(self, label: ForgeLabel, expected_phase: str): """Each label maps to the expected workflow phase.""" labels = ["forge:managed", label.value] diff --git a/tests/flows/status_transitions/test_plan_rejected.py b/tests/flows/status_transitions/test_plan_rejected.py index ddd6e13d..8dbc7d60 100644 --- a/tests/flows/status_transitions/test_plan_rejected.py +++ b/tests/flows/status_transitions/test_plan_rejected.py @@ -1,11 +1,10 @@ """Tests for Plan rejection and revision cycles.""" - import pytest from forge.models.workflow import TicketType -from forge.workflow.gates import route_plan_approval from forge.workflow.feature.state import create_initial_feature_state as create_initial_state +from forge.workflow.gates import route_plan_approval class TestPlanRejectedFullRegen: @@ -26,7 +25,8 @@ def plan_pending_state(self): state["epic_keys"] = ["TEST-124", "TEST-125", "TEST-126"] return state - def test_feature_level_rejection_regenerates_all(self, plan_pending_state): + @pytest.mark.asyncio + async def test_feature_level_rejection_regenerates_all(self, plan_pending_state): """Feature-level rejection regenerates all epics.""" plan_pending_state["context"] = { "labels": ["forge:managed", "forge:plan-pending"], @@ -35,11 +35,12 @@ def test_feature_level_rejection_regenerates_all(self, plan_pending_state): plan_pending_state["feedback_comment"] = "The entire breakdown is wrong. Start over." plan_pending_state["revision_requested"] = True - result = route_plan_approval(plan_pending_state) + result = await route_plan_approval(plan_pending_state) assert result == "regenerate_all_epics" - def test_all_epics_will_be_deleted(self, plan_pending_state): + @pytest.mark.asyncio + async def test_all_epics_will_be_deleted(self, plan_pending_state): """Full regeneration implies all existing epics deleted.""" plan_pending_state["context"] = { "labels": ["forge:managed", "forge:plan-pending"], @@ -51,7 +52,7 @@ def test_all_epics_will_be_deleted(self, plan_pending_state): # Verify all epic keys exist before regeneration decision assert len(plan_pending_state["epic_keys"]) == 3 - result = route_plan_approval(plan_pending_state) + result = await route_plan_approval(plan_pending_state) assert result == "regenerate_all_epics" @@ -72,7 +73,8 @@ def plan_with_epic_issue(self): state["current_epic_key"] = "TEST-125" # The problematic epic return state - def test_single_epic_rejection_updates_only_that_epic(self, plan_with_epic_issue): + @pytest.mark.asyncio + async def test_single_epic_rejection_updates_only_that_epic(self, plan_with_epic_issue): """Single epic rejection only updates that epic.""" plan_with_epic_issue["context"] = { "labels": ["forge:managed", "forge:plan-pending"], @@ -82,11 +84,12 @@ def test_single_epic_rejection_updates_only_that_epic(self, plan_with_epic_issue plan_with_epic_issue["feedback_comment"] = "Epic 2 scope is too narrow." plan_with_epic_issue["revision_requested"] = True - result = route_plan_approval(plan_with_epic_issue) + result = await route_plan_approval(plan_with_epic_issue) assert result == "update_single_epic" - def test_other_epics_preserved(self, plan_with_epic_issue): + @pytest.mark.asyncio + async def test_other_epics_preserved(self, plan_with_epic_issue): """Other epics are preserved when one is revised.""" plan_with_epic_issue["context"] = { "labels": ["forge:managed", "forge:plan-pending"], @@ -100,7 +103,7 @@ def test_other_epics_preserved(self, plan_with_epic_issue): assert "TEST-124" in plan_with_epic_issue["epic_keys"] assert "TEST-126" in plan_with_epic_issue["epic_keys"] - result = route_plan_approval(plan_with_epic_issue) + result = await route_plan_approval(plan_with_epic_issue) assert result == "update_single_epic" @@ -120,7 +123,8 @@ def plan_partial_approval(self): state["epic_keys"] = ["TEST-124", "TEST-125", "TEST-126"] return state - def test_some_approved_one_rejected(self, plan_partial_approval): + @pytest.mark.asyncio + async def test_some_approved_one_rejected(self, plan_partial_approval): """Some epics approved, one needs revision.""" plan_partial_approval["context"] = { "labels": ["forge:managed", "forge:plan-pending"], @@ -132,16 +136,17 @@ def test_some_approved_one_rejected(self, plan_partial_approval): plan_partial_approval["feedback_comment"] = "Epic 3 needs more detail." plan_partial_approval["revision_requested"] = True - result = route_plan_approval(plan_partial_approval) + result = await route_plan_approval(plan_partial_approval) assert result == "update_single_epic" - def test_all_approved_routes_to_tasks(self, plan_partial_approval): + @pytest.mark.asyncio + async def test_all_approved_routes_to_tasks(self, plan_partial_approval): """All epics approved routes to task generation when resumed.""" # Workflow is resumed from pause on approval webhook plan_partial_approval["is_paused"] = False - result = route_plan_approval(plan_partial_approval) + result = await route_plan_approval(plan_partial_approval) assert result == "generate_tasks" @@ -170,12 +175,13 @@ def plan_with_spec_issue(self): state["revision_requested"] = True return state - def test_spec_scope_feedback_noted(self, plan_with_spec_issue): + @pytest.mark.asyncio + async def test_spec_scope_feedback_noted(self, plan_with_spec_issue): """Feedback targeting spec is noted but routes to plan regen.""" # The feedback mentions spec issues assert "Spec" in plan_with_spec_issue["feedback_comment"] - result = route_plan_approval(plan_with_spec_issue) + result = await route_plan_approval(plan_with_spec_issue) # Currently routes to regenerate (future: could escalate) assert result in ["regenerate_all_epics", "update_single_epic"] diff --git a/tests/flows/status_transitions/test_spec_rejected.py b/tests/flows/status_transitions/test_spec_rejected.py index 59e577ac..c7caf043 100644 --- a/tests/flows/status_transitions/test_spec_rejected.py +++ b/tests/flows/status_transitions/test_spec_rejected.py @@ -1,11 +1,10 @@ """Tests for Spec rejection and revision cycles.""" - import pytest from forge.models.workflow import TicketType -from forge.workflow.gates import route_spec_approval from forge.workflow.feature.state import create_initial_feature_state as create_initial_state +from forge.workflow.gates import route_spec_approval class TestSpecRejectedOnce: 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 ba2d17b4..22662233 100644 --- a/tests/integration/orchestrator/test_ci_fix_attempt_status_comments.py +++ b/tests/integration/orchestrator/test_ci_fix_attempt_status_comments.py @@ -68,23 +68,36 @@ async def test_first_attempt_posts_comment_with_1_of_max(self): state["ci_fix_max_attempts"] = 3 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.prepare_workspace") as mock_prepare: + 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.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()): - with patch("forge.workflow.nodes.ci_evaluator._collect_error_info", return_value="errors"): - with patch("forge.workflow.nodes.ci_evaluator.load_prompt", return_value="prompt"): - with patch("pathlib.Path.mkdir"): - with patch("pathlib.Path.write_text"): - with patch("pathlib.Path.exists", return_value=False): - await attempt_ci_fix(state) + with patch( + "forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", + AsyncMock(), + ), patch( + "forge.workflow.nodes.ci_evaluator._collect_error_info", + return_value="errors", + ), patch( + "forge.workflow.nodes.ci_evaluator.load_prompt", + return_value="prompt", + ), patch("pathlib.Path.mkdir"), patch("pathlib.Path.write_text"): + with patch("pathlib.Path.exists", return_value=False): + await attempt_ci_fix(state) # Verify status comment posted with correct format assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args assert comment_call[0][0] == "FEAT-300" - assert comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (1/3)." + assert ( + comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (1/3)." + ) # Verify JiraClient closed assert mock_jira.close.call_count == 1 @@ -115,23 +128,36 @@ async def test_second_attempt_posts_comment_with_2_of_max(self): state["ci_fix_max_attempts"] = 3 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.prepare_workspace") as mock_prepare: + 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.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()): - with patch("forge.workflow.nodes.ci_evaluator._collect_error_info", return_value="errors"): - with patch("forge.workflow.nodes.ci_evaluator.load_prompt", return_value="prompt"): - with patch("pathlib.Path.mkdir"): - with patch("pathlib.Path.write_text"): - with patch("pathlib.Path.exists", return_value=False): - await attempt_ci_fix(state) + with patch( + "forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", + AsyncMock(), + ), patch( + "forge.workflow.nodes.ci_evaluator._collect_error_info", + return_value="errors", + ), patch( + "forge.workflow.nodes.ci_evaluator.load_prompt", + return_value="prompt", + ), patch("pathlib.Path.mkdir"), patch("pathlib.Path.write_text"): + with patch("pathlib.Path.exists", return_value=False): + await attempt_ci_fix(state) # Verify status comment posted with correct format assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args assert comment_call[0][0] == "FEAT-301" - assert comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (2/3)." + assert ( + comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (2/3)." + ) @pytest.mark.asyncio async def test_final_attempt_posts_comment_with_max_of_max(self): @@ -159,23 +185,36 @@ async def test_final_attempt_posts_comment_with_max_of_max(self): state["ci_fix_max_attempts"] = 3 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.prepare_workspace") as mock_prepare: + 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.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()): - with patch("forge.workflow.nodes.ci_evaluator._collect_error_info", return_value="errors"): - with patch("forge.workflow.nodes.ci_evaluator.load_prompt", return_value="prompt"): - with patch("pathlib.Path.mkdir"): - with patch("pathlib.Path.write_text"): - with patch("pathlib.Path.exists", return_value=False): - await attempt_ci_fix(state) + with patch( + "forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", + AsyncMock(), + ), patch( + "forge.workflow.nodes.ci_evaluator._collect_error_info", + return_value="errors", + ), patch( + "forge.workflow.nodes.ci_evaluator.load_prompt", + return_value="prompt", + ), patch("pathlib.Path.mkdir"), patch("pathlib.Path.write_text"): + with patch("pathlib.Path.exists", return_value=False): + await attempt_ci_fix(state) # Verify status comment posted with correct format assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args assert comment_call[0][0] == "FEAT-302" - assert comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (3/3)." + assert ( + comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (3/3)." + ) @pytest.mark.asyncio async def test_comment_posted_to_feature_ticket_not_task(self): @@ -203,17 +242,28 @@ async def test_comment_posted_to_feature_ticket_not_task(self): state["ci_fix_max_attempts"] = 5 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.prepare_workspace") as mock_prepare: + 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.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()): - with patch("forge.workflow.nodes.ci_evaluator._collect_error_info", return_value="errors"): - with patch("forge.workflow.nodes.ci_evaluator.load_prompt", return_value="prompt"): - with patch("pathlib.Path.mkdir"): - with patch("pathlib.Path.write_text"): - with patch("pathlib.Path.exists", return_value=False): - await attempt_ci_fix(state) + with patch( + "forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", + AsyncMock(), + ), patch( + "forge.workflow.nodes.ci_evaluator._collect_error_info", + return_value="errors", + ), patch( + "forge.workflow.nodes.ci_evaluator.load_prompt", + return_value="prompt", + ), patch("pathlib.Path.mkdir"), patch("pathlib.Path.write_text"): + with patch("pathlib.Path.exists", return_value=False): + await attempt_ci_fix(state) # Verify comment posted to feature ticket (FEAT-303), not task tickets (TASK-001, TASK-002) assert mock_jira.add_comment.call_count == 1 @@ -235,10 +285,10 @@ async def test_multiple_attempts_show_incrementing_counts(self): # Collect all comments posted comments = [] - + def capture_comment(ticket_key, message): comments.append((ticket_key, message)) - + mock_jira.add_comment.side_effect = capture_comment base_state = create_initial_feature_state( @@ -261,19 +311,31 @@ def capture_comment(ticket_key, message): # Simulate three attempts for attempt in [1, 2, 3]: state = {**base_state, "ci_fix_attempt": attempt} - + 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.prepare_workspace") as mock_prepare: + 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.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()): - with patch("forge.workflow.nodes.ci_evaluator._collect_error_info", return_value="errors"): - with patch("forge.workflow.nodes.ci_evaluator.load_prompt", return_value="prompt"): - with patch("pathlib.Path.mkdir"): - with patch("pathlib.Path.write_text"): - with patch("pathlib.Path.exists", return_value=False): - await attempt_ci_fix(state) + with patch( + "forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", + AsyncMock(), + ), patch( + "forge.workflow.nodes.ci_evaluator._collect_error_info", + return_value="errors", + ), patch( + "forge.workflow.nodes.ci_evaluator.load_prompt", + return_value="prompt", + ), patch("pathlib.Path.mkdir"), patch("pathlib.Path.write_text"), patch( + "pathlib.Path.exists", return_value=False + ): + await attempt_ci_fix(state) # Verify three comments posted with correct counts assert len(comments) == 3 @@ -307,22 +369,35 @@ async def test_different_max_attempts_values(self): state["ci_fix_max_attempts"] = 5 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.prepare_workspace") as mock_prepare: + 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.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()): - with patch("forge.workflow.nodes.ci_evaluator._collect_error_info", return_value="errors"): - with patch("forge.workflow.nodes.ci_evaluator.load_prompt", return_value="prompt"): - with patch("pathlib.Path.mkdir"): - with patch("pathlib.Path.write_text"): - with patch("pathlib.Path.exists", return_value=False): - await attempt_ci_fix(state) + with patch( + "forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", + AsyncMock(), + ), patch( + "forge.workflow.nodes.ci_evaluator._collect_error_info", + return_value="errors", + ), patch( + "forge.workflow.nodes.ci_evaluator.load_prompt", + return_value="prompt", + ), patch("pathlib.Path.mkdir"), patch("pathlib.Path.write_text"): + with patch("pathlib.Path.exists", return_value=False): + await attempt_ci_fix(state) # Verify comment uses max_attempts=5 assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args - assert comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (2/5)." + assert ( + comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (2/5)." + ) class TestCIFixAttemptErrorHandling: @@ -334,7 +409,7 @@ async def test_workflow_continues_when_comment_posting_fails(self, caplog): mock_jira = create_mock_jira_client() # Simulate comment posting failure mock_jira.add_comment.side_effect = Exception("Jira API error") - + mock_runner = create_mock_container_runner() mock_github = create_mock_github_client() @@ -357,17 +432,28 @@ async def test_workflow_continues_when_comment_posting_fails(self, caplog): state["ci_fix_max_attempts"] = 3 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.prepare_workspace") as mock_prepare: + 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.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()): - with patch("forge.workflow.nodes.ci_evaluator._collect_error_info", return_value="errors"): - with patch("forge.workflow.nodes.ci_evaluator.load_prompt", return_value="prompt"): - with patch("pathlib.Path.mkdir"): - with patch("pathlib.Path.write_text"): - with patch("pathlib.Path.exists", return_value=False): - result = await attempt_ci_fix(state) + with patch( + "forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", + AsyncMock(), + ), patch( + "forge.workflow.nodes.ci_evaluator._collect_error_info", + return_value="errors", + ), patch( + "forge.workflow.nodes.ci_evaluator.load_prompt", + return_value="prompt", + ), patch("pathlib.Path.mkdir"), patch("pathlib.Path.write_text"): + with patch("pathlib.Path.exists", return_value=False): + result = await attempt_ci_fix(state) # Verify workflow continues (doesn't raise exception) assert result is not None @@ -380,7 +466,7 @@ async def test_jira_client_closed_even_on_comment_error(self): mock_jira = create_mock_jira_client() # Simulate comment posting failure mock_jira.add_comment.side_effect = Exception("Jira API error") - + mock_runner = create_mock_container_runner() mock_github = create_mock_github_client() @@ -403,17 +489,28 @@ async def test_jira_client_closed_even_on_comment_error(self): state["ci_fix_max_attempts"] = 3 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.prepare_workspace") as mock_prepare: + 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.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()): - with patch("forge.workflow.nodes.ci_evaluator._collect_error_info", return_value="errors"): - with patch("forge.workflow.nodes.ci_evaluator.load_prompt", return_value="prompt"): - with patch("pathlib.Path.mkdir"): - with patch("pathlib.Path.write_text"): - with patch("pathlib.Path.exists", return_value=False): - await attempt_ci_fix(state) + with patch( + "forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", + AsyncMock(), + ), patch( + "forge.workflow.nodes.ci_evaluator._collect_error_info", + return_value="errors", + ), patch( + "forge.workflow.nodes.ci_evaluator.load_prompt", + return_value="prompt", + ), patch("pathlib.Path.mkdir"), patch("pathlib.Path.write_text"): + with patch("pathlib.Path.exists", return_value=False): + await attempt_ci_fix(state) # Verify JiraClient closed despite error assert mock_jira.close.call_count == 1 diff --git a/tests/integration/orchestrator/test_pr_creation_status_comments.py b/tests/integration/orchestrator/test_pr_creation_status_comments.py index a7fb1ea4..f7de43f8 100644 --- a/tests/integration/orchestrator/test_pr_creation_status_comments.py +++ b/tests/integration/orchestrator/test_pr_creation_status_comments.py @@ -50,7 +50,10 @@ async def test_pr_creation_posts_comment_with_pr_number(self): assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args assert comment_call[0][0] == "FEAT-200" - assert comment_call[0][1] == "🚀 Pull request #123 created and submitted. Waiting for CI checks to complete." + assert ( + comment_call[0][1] + == "🚀 Pull request #123 created and submitted. Waiting for CI checks to complete." + ) # Verify workflow paused assert result["is_paused"] is True @@ -100,6 +103,7 @@ async def test_pr_creation_adds_ci_pending_label(self): assert label_call[0][0] == "FEAT-200" # Check that it's the CI_PENDING label (value is "forge:ci-pending") from forge.models.workflow import ForgeLabel + assert label_call[0][1] == ForgeLabel.TASK_CI_PENDING @pytest.mark.asyncio @@ -146,7 +150,10 @@ async def test_pr_creation_posts_comment_without_pr_number(self): assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args assert comment_call[0][0] == "FEAT-201" - assert comment_call[0][1] == "🚀 Pull request created and submitted. Waiting for CI checks to complete." + assert ( + comment_call[0][1] + == "🚀 Pull request created and submitted. Waiting for CI checks to complete." + ) # Verify workflow paused assert result["is_paused"] is True @@ -225,7 +232,9 @@ async def test_workflow_continues_when_label_removal_fails(self, caplog): assert result["current_node"] == "wait_for_ci_gate" # Verify error logged - assert any("Failed to remove implementing label" in record.message for record in caplog.records) + assert any( + "Failed to remove implementing label" in record.message for record in caplog.records + ) @pytest.mark.asyncio async def test_workflow_continues_when_label_setting_fails(self, caplog): diff --git a/tests/integration/orchestrator/test_workflow_execution.py b/tests/integration/orchestrator/test_workflow_execution.py index 3db1ab39..88ed75bc 100644 --- a/tests/integration/orchestrator/test_workflow_execution.py +++ b/tests/integration/orchestrator/test_workflow_execution.py @@ -158,9 +158,10 @@ async def test_feature_runs_through_prd_and_pauses( ) # Mock external dependencies - with patch("forge.workflow.nodes.prd_generation.JiraClient") as MockJira, \ - patch("forge.workflow.nodes.prd_generation.ForgeAgent") as MockAgent: - + with ( + patch("forge.workflow.nodes.prd_generation.JiraClient") as MockJira, + patch("forge.workflow.nodes.prd_generation.ForgeAgent") as MockAgent, + ): MockJira.return_value = mock_jira_client MockAgent.return_value = mock_agent @@ -195,9 +196,10 @@ async def test_workflow_state_persisted_via_checkpointer( ticket_type=TicketType.FEATURE, ) - with patch("forge.workflow.nodes.prd_generation.JiraClient") as MockJira, \ - patch("forge.workflow.nodes.prd_generation.ForgeAgent") as MockAgent: - + with ( + patch("forge.workflow.nodes.prd_generation.JiraClient") as MockJira, + patch("forge.workflow.nodes.prd_generation.ForgeAgent") as MockAgent, + ): MockJira.return_value = mock_jira_client MockAgent.return_value = mock_agent @@ -224,6 +226,7 @@ async def test_bug_runs_through_rca_and_pauses( """Bug workflow should generate RCA and pause at approval gate.""" # Update mock for bug issue from forge.integrations.jira.models import JiraIssue + mock_jira_client.get_issue = AsyncMock( return_value=JiraIssue( key="BUG-456", @@ -245,10 +248,11 @@ async def test_bug_runs_through_rca_and_pauses( ticket_type=TicketType.BUG, ) - with patch("forge.workflow.nodes.bug_workflow.JiraClient") as MockJira, \ - patch("forge.workflow.nodes.bug_workflow.ForgeAgent") as MockAgent, \ - patch("forge.workflow.nodes.bug_workflow.get_settings") as mock_settings: - + with ( + patch("forge.workflow.nodes.bug_workflow.JiraClient") as MockJira, + patch("forge.workflow.nodes.bug_workflow.ForgeAgent") as MockAgent, + patch("forge.workflow.nodes.bug_workflow.get_settings") as mock_settings, + ): MockJira.return_value = mock_jira_client MockAgent.return_value = mock_agent mock_settings.return_value = MagicMock() @@ -282,9 +286,10 @@ async def test_workflow_resumes_from_checkpoint( ticket_type=TicketType.FEATURE, ) - with patch("forge.workflow.nodes.prd_generation.JiraClient") as MockJira, \ - patch("forge.workflow.nodes.prd_generation.ForgeAgent") as MockAgent: - + with ( + patch("forge.workflow.nodes.prd_generation.JiraClient") as MockJira, + patch("forge.workflow.nodes.prd_generation.ForgeAgent") as MockAgent, + ): MockJira.return_value = mock_jira_client MockAgent.return_value = mock_agent diff --git a/tests/integration/workflow/test_pr_ci_status_updates.py b/tests/integration/workflow/test_pr_ci_status_updates.py index e6cde416..c461d98a 100644 --- a/tests/integration/workflow/test_pr_ci_status_updates.py +++ b/tests/integration/workflow/test_pr_ci_status_updates.py @@ -22,7 +22,7 @@ def create_mock_jira_client(): """Create a mock JiraClient with required methods for testing. - + Returns: MagicMock: Mock JiraClient with async methods for comment posting and label management. """ @@ -36,7 +36,7 @@ def create_mock_jira_client(): def create_mock_container_runner(): """Create a mock ContainerRunner that succeeds. - + Returns: MagicMock: Mock ContainerRunner with async run method. """ @@ -47,7 +47,7 @@ def create_mock_container_runner(): def create_mock_github_client(): """Create a mock GitHubClient. - + Returns: MagicMock: Mock GitHubClient with async close method. """ @@ -62,7 +62,7 @@ class TestPRCreationWithPRNumber: @pytest.mark.asyncio async def test_pr_creation_posts_comment_with_pr_number(self): """TS-006: Verify comment posted with PR number when available. - + This test ensures that when a PR is created successfully with a valid PR number, the status comment includes the PR number in the expected format. """ @@ -83,7 +83,10 @@ async def test_pr_creation_posts_comment_with_pr_number(self): assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args assert comment_call[0][0] == "FEAT-200" - assert comment_call[0][1] == "🚀 Pull request #123 created and submitted. Waiting for CI checks to complete." + assert ( + comment_call[0][1] + == "🚀 Pull request #123 created and submitted. Waiting for CI checks to complete." + ) # Verify workflow paused assert result["is_paused"] is True @@ -92,7 +95,7 @@ async def test_pr_creation_posts_comment_with_pr_number(self): @pytest.mark.asyncio async def test_pr_creation_removes_implementing_label(self): """TS-006: Verify forge:implementing label removed from feature ticket. - + This test ensures the label transition removes the implementing label when PR creation occurs. """ @@ -118,7 +121,7 @@ async def test_pr_creation_removes_implementing_label(self): @pytest.mark.asyncio async def test_pr_creation_adds_ci_pending_label(self): """TS-006: Verify forge:ci-pending label added to feature ticket. - + This test ensures the label transition adds the ci-pending label when PR creation occurs. """ @@ -141,12 +144,13 @@ async def test_pr_creation_adds_ci_pending_label(self): assert label_call[0][0] == "FEAT-200" # Check that it's the CI_PENDING label (value is "forge:ci-pending") from forge.models.workflow import ForgeLabel + assert label_call[0][1] == ForgeLabel.TASK_CI_PENDING @pytest.mark.asyncio async def test_pr_creation_jira_client_properly_closed(self): """TS-006: Verify JiraClient properly closed after operations. - + This test ensures proper resource cleanup by verifying the JiraClient is closed in the finally block. """ @@ -173,7 +177,7 @@ class TestCIFixAttemptStatusComments: @pytest.mark.asyncio async def test_first_attempt_posts_comment_with_1_of_3(self): """TS-007: Verify first CI fix attempt posts comment with '1/3' format. - + This test ensures the first fix attempt shows the correct count format. """ mock_jira = create_mock_jira_client() @@ -199,23 +203,36 @@ async def test_first_attempt_posts_comment_with_1_of_3(self): state["ci_fix_max_attempts"] = 3 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.prepare_workspace") as mock_prepare: + 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.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()): - with patch("forge.workflow.nodes.ci_evaluator._collect_error_info", return_value="errors"): - with patch("forge.workflow.nodes.ci_evaluator.load_prompt", return_value="prompt"): - with patch("pathlib.Path.mkdir"): - with patch("pathlib.Path.write_text"): - with patch("pathlib.Path.exists", return_value=False): - await attempt_ci_fix(state) + with patch( + "forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", + AsyncMock(), + ), patch( + "forge.workflow.nodes.ci_evaluator._collect_error_info", + return_value="errors", + ), patch( + "forge.workflow.nodes.ci_evaluator.load_prompt", + return_value="prompt", + ), patch("pathlib.Path.mkdir"), patch("pathlib.Path.write_text"): + with patch("pathlib.Path.exists", return_value=False): + await attempt_ci_fix(state) # Verify status comment posted with correct format "1/3" assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args assert comment_call[0][0] == "FEAT-300" - assert comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (1/3)." + assert ( + comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (1/3)." + ) # Verify JiraClient closed assert mock_jira.close.call_count == 1 @@ -223,7 +240,7 @@ async def test_first_attempt_posts_comment_with_1_of_3(self): @pytest.mark.asyncio async def test_second_attempt_posts_comment_with_2_of_3(self): """TS-007: Verify second CI fix attempt posts comment with '2/3' format. - + This test ensures the second fix attempt shows the correct count format. """ mock_jira = create_mock_jira_client() @@ -249,28 +266,41 @@ async def test_second_attempt_posts_comment_with_2_of_3(self): state["ci_fix_max_attempts"] = 3 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.prepare_workspace") as mock_prepare: + 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.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()): - with patch("forge.workflow.nodes.ci_evaluator._collect_error_info", return_value="errors"): - with patch("forge.workflow.nodes.ci_evaluator.load_prompt", return_value="prompt"): - with patch("pathlib.Path.mkdir"): - with patch("pathlib.Path.write_text"): - with patch("pathlib.Path.exists", return_value=False): - await attempt_ci_fix(state) + with patch( + "forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", + AsyncMock(), + ), patch( + "forge.workflow.nodes.ci_evaluator._collect_error_info", + return_value="errors", + ), patch( + "forge.workflow.nodes.ci_evaluator.load_prompt", + return_value="prompt", + ), patch("pathlib.Path.mkdir"), patch("pathlib.Path.write_text"): + with patch("pathlib.Path.exists", return_value=False): + await attempt_ci_fix(state) # Verify status comment posted with correct format "2/3" assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args assert comment_call[0][0] == "FEAT-301" - assert comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (2/3)." + assert ( + comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (2/3)." + ) @pytest.mark.asyncio async def test_third_attempt_posts_comment_with_3_of_3(self): """TS-007: Verify third CI fix attempt posts comment with '3/3' format. - + This test ensures the final fix attempt shows the correct count format. """ mock_jira = create_mock_jira_client() @@ -296,23 +326,36 @@ async def test_third_attempt_posts_comment_with_3_of_3(self): state["ci_fix_max_attempts"] = 3 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.prepare_workspace") as mock_prepare: + 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.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()): - with patch("forge.workflow.nodes.ci_evaluator._collect_error_info", return_value="errors"): - with patch("forge.workflow.nodes.ci_evaluator.load_prompt", return_value="prompt"): - with patch("pathlib.Path.mkdir"): - with patch("pathlib.Path.write_text"): - with patch("pathlib.Path.exists", return_value=False): - await attempt_ci_fix(state) + with patch( + "forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", + AsyncMock(), + ), patch( + "forge.workflow.nodes.ci_evaluator._collect_error_info", + return_value="errors", + ), patch( + "forge.workflow.nodes.ci_evaluator.load_prompt", + return_value="prompt", + ), patch("pathlib.Path.mkdir"), patch("pathlib.Path.write_text"): + with patch("pathlib.Path.exists", return_value=False): + await attempt_ci_fix(state) # Verify status comment posted with correct format "3/3" assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args assert comment_call[0][0] == "FEAT-302" - assert comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (3/3)." + assert ( + comment_call[0][1] == "🔧 CI checks failed. Analyzing failure and attempting fix (3/3)." + ) class TestPRCreationFallbackWithoutPRNumber: @@ -321,7 +364,7 @@ class TestPRCreationFallbackWithoutPRNumber: @pytest.mark.asyncio async def test_pr_creation_posts_fallback_comment_without_pr_number(self): """TS-014: Verify fallback comment posted when PR number unavailable. - + This test ensures that when GitHub PR creation doesn't return a PR number, the fallback comment text is used instead of including a null/missing number. """ @@ -343,7 +386,10 @@ async def test_pr_creation_posts_fallback_comment_without_pr_number(self): assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args assert comment_call[0][0] == "FEAT-201" - assert comment_call[0][1] == "🚀 Pull request created and submitted. Waiting for CI checks to complete." + assert ( + comment_call[0][1] + == "🚀 Pull request created and submitted. Waiting for CI checks to complete." + ) # Verify workflow still paused correctly assert result["is_paused"] is True @@ -352,7 +398,7 @@ async def test_pr_creation_posts_fallback_comment_without_pr_number(self): @pytest.mark.asyncio async def test_pr_creation_without_pr_number_still_updates_labels(self): """TS-014: Verify label transitions still occur when PR number unavailable. - + This test ensures that missing PR number doesn't prevent label transitions from occurring correctly. """ @@ -381,6 +427,7 @@ async def test_pr_creation_without_pr_number_still_updates_labels(self): label_call = mock_jira.set_workflow_label.call_args assert label_call[0][0] == "FEAT-202" from forge.models.workflow import ForgeLabel + assert label_call[0][1] == ForgeLabel.TASK_CI_PENDING @@ -390,7 +437,7 @@ class TestErrorHandling: @pytest.mark.asyncio async def test_workflow_continues_when_pr_comment_posting_fails(self, caplog): """Verify workflow continues when PR creation comment posting fails. - + This test ensures that Jira API failures don't block the workflow from continuing to the next state. """ @@ -419,7 +466,7 @@ async def test_workflow_continues_when_pr_comment_posting_fails(self, caplog): @pytest.mark.asyncio async def test_workflow_continues_when_label_removal_fails(self, caplog): """Verify workflow continues when label removal fails. - + This test ensures that label API failures are properly suppressed and logged. """ mock_jira = create_mock_jira_client() @@ -447,7 +494,7 @@ async def test_workflow_continues_when_label_removal_fails(self, caplog): @pytest.mark.asyncio async def test_workflow_continues_when_ci_attempt_comment_posting_fails(self, caplog): """Verify workflow continues when CI attempt comment posting fails. - + This test ensures that Jira failures during CI fix attempts don't block the workflow from continuing. """ @@ -476,17 +523,28 @@ async def test_workflow_continues_when_ci_attempt_comment_posting_fails(self, ca state["ci_fix_max_attempts"] = 3 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.prepare_workspace") as mock_prepare: + 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.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()): - with patch("forge.workflow.nodes.ci_evaluator._collect_error_info", return_value="errors"): - with patch("forge.workflow.nodes.ci_evaluator.load_prompt", return_value="prompt"): - with patch("pathlib.Path.mkdir"): - with patch("pathlib.Path.write_text"): - with patch("pathlib.Path.exists", return_value=False): - result = await attempt_ci_fix(state) + with patch( + "forge.workflow.nodes.ci_evaluator._fetch_ci_logs_and_artifacts", + AsyncMock(), + ), patch( + "forge.workflow.nodes.ci_evaluator._collect_error_info", + return_value="errors", + ), patch( + "forge.workflow.nodes.ci_evaluator.load_prompt", + return_value="prompt", + ), patch("pathlib.Path.mkdir"), patch("pathlib.Path.write_text"): + with patch("pathlib.Path.exists", return_value=False): + result = await attempt_ci_fix(state) # Verify workflow continues despite failure assert "next_node" in result or "error" in result or result is not None diff --git a/tests/sandbox/test_task_execution.py b/tests/sandbox/test_task_execution.py index ca04e4f7..bc2b2d61 100644 --- a/tests/sandbox/test_task_execution.py +++ b/tests/sandbox/test_task_execution.py @@ -250,7 +250,6 @@ async def test_build_and_test_recovery_workflow_iterative_self_correction( assert state_after_success["commit_info"]["committed"] is True assert state_after_success["commit_info"]["sha"] == "abcdef1234567890" - @pytest.mark.asyncio @patch("forge.workflow.nodes.workspace_setup.get_workspace_manager") async def test_teardown_workspace_secure_destruction(self, mock_get_manager: MagicMock) -> None: diff --git a/tests/test_sandbox_runner.py b/tests/test_sandbox_runner.py index e4e02c24..76530a14 100644 --- a/tests/test_sandbox_runner.py +++ b/tests/test_sandbox_runner.py @@ -21,6 +21,7 @@ def test_runner_init(self): def test_podman_exists(self): """Test podman is available.""" import shutil + assert shutil.which("podman") is not None @pytest.mark.asyncio @@ -46,10 +47,14 @@ async def test_simple_container_run(self): result = subprocess.run( [ - "podman", "run", "--rm", - "-v", f"{workspace}:/workspace:Z", + "podman", + "run", + "--rm", + "-v", + f"{workspace}:/workspace:Z", "alpine:latest", - "cat", "/workspace/test.txt", + "cat", + "/workspace/test.txt", ], capture_output=True, text=True, diff --git a/tests/unit/api/routes/test_github_webhook.py b/tests/unit/api/routes/test_github_webhook.py index 2acf5d62..54b9aa6b 100644 --- a/tests/unit/api/routes/test_github_webhook.py +++ b/tests/unit/api/routes/test_github_webhook.py @@ -8,14 +8,14 @@ import pytest from httpx import ASGITransport, AsyncClient from pydantic import SecretStr + +from forge.main import app from tests.fixtures.github_payloads import ( WEBHOOK_CHECK_RUN_COMPLETED_FAILURE, WEBHOOK_CHECK_RUN_COMPLETED_SUCCESS, WEBHOOK_PULL_REQUEST_REVIEW_APPROVED, ) -from forge.main import app - def compute_signature(payload: bytes, secret: str) -> str: """Compute GitHub webhook signature with sha256= prefix.""" @@ -46,8 +46,7 @@ async def test_valid_webhook_returns_202(self): with patch("forge.api.routes.github.get_settings", return_value=mock_settings): with patch("forge.api.routes.github.QueueProducer", return_value=mock_producer): async with AsyncClient( - transport=ASGITransport(app=app), - base_url="http://test" + transport=ASGITransport(app=app), base_url="http://test" ) as client: response = await client.post( "/api/v1/webhooks/github", @@ -72,8 +71,7 @@ async def test_invalid_signature_returns_401(self): with patch("forge.api.routes.github.get_settings", return_value=mock_settings): async with AsyncClient( - transport=ASGITransport(app=app), - base_url="http://test" + transport=ASGITransport(app=app), base_url="http://test" ) as client: response = await client.post( "/api/v1/webhooks/github", @@ -97,8 +95,7 @@ async def test_missing_signature_returns_401(self): with patch("forge.api.routes.github.get_settings", return_value=mock_settings): async with AsyncClient( - transport=ASGITransport(app=app), - base_url="http://test" + transport=ASGITransport(app=app), base_url="http://test" ) as client: response = await client.post( "/api/v1/webhooks/github", @@ -127,8 +124,7 @@ async def test_check_run_success_published(self): with patch("forge.api.routes.github.get_settings", return_value=mock_settings): with patch("forge.api.routes.github.QueueProducer", return_value=mock_producer): async with AsyncClient( - transport=ASGITransport(app=app), - base_url="http://test" + transport=ASGITransport(app=app), base_url="http://test" ) as client: response = await client.post( "/api/v1/webhooks/github", @@ -160,8 +156,7 @@ async def test_check_run_failure_published(self): with patch("forge.api.routes.github.get_settings", return_value=mock_settings): with patch("forge.api.routes.github.QueueProducer", return_value=mock_producer): async with AsyncClient( - transport=ASGITransport(app=app), - base_url="http://test" + transport=ASGITransport(app=app), base_url="http://test" ) as client: response = await client.post( "/api/v1/webhooks/github", @@ -193,8 +188,7 @@ async def test_pr_review_approved_published(self): with patch("forge.api.routes.github.get_settings", return_value=mock_settings): with patch("forge.api.routes.github.QueueProducer", return_value=mock_producer): async with AsyncClient( - transport=ASGITransport(app=app), - base_url="http://test" + transport=ASGITransport(app=app), base_url="http://test" ) as client: response = await client.post( "/api/v1/webhooks/github", @@ -224,8 +218,12 @@ def test_extract_check_conclusion(self): """Extract check run conclusion.""" from forge.integrations.github.webhooks import parse_github_webhook - 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") + 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 success_data.check_conclusion == "success" assert failure_data.check_conclusion == "failure" diff --git a/tests/unit/api/routes/test_health.py b/tests/unit/api/routes/test_health.py index 79d94dc7..fc9b259c 100644 --- a/tests/unit/api/routes/test_health.py +++ b/tests/unit/api/routes/test_health.py @@ -20,8 +20,7 @@ async def test_health_returns_200(self): with patch("forge.api.routes.health.get_redis_client", return_value=mock_redis): async with AsyncClient( - transport=ASGITransport(app=app), - base_url="http://test" + transport=ASGITransport(app=app), base_url="http://test" ) as client: response = await client.get("/api/v1/health") @@ -38,8 +37,7 @@ async def test_health_includes_version(self): with patch("forge.api.routes.health.get_redis_client", return_value=mock_redis): async with AsyncClient( - transport=ASGITransport(app=app), - base_url="http://test" + transport=ASGITransport(app=app), base_url="http://test" ) as client: response = await client.get("/api/v1/health") @@ -53,10 +51,7 @@ class TestReadinessEndpoint: @pytest.mark.asyncio async def test_ready_with_healthy_dependencies(self): """Ready returns 200 (always ready in current impl).""" - async with AsyncClient( - transport=ASGITransport(app=app), - base_url="http://test" - ) as client: + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: response = await client.get("/api/v1/ready") assert response.status_code == 200 @@ -67,10 +62,7 @@ async def test_ready_with_healthy_dependencies(self): async def test_ready_with_unhealthy_redis(self): """Ready endpoint doesn't check Redis (always returns ready).""" # Current implementation doesn't check Redis for readiness - async with AsyncClient( - transport=ASGITransport(app=app), - base_url="http://test" - ) as client: + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: response = await client.get("/api/v1/ready") assert response.status_code == 200 @@ -84,10 +76,7 @@ class TestLivenessEndpoint: @pytest.mark.asyncio async def test_live_returns_200(self): """Liveness endpoint always returns 200.""" - async with AsyncClient( - transport=ASGITransport(app=app), - base_url="http://test" - ) as client: + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: response = await client.get("/api/v1/live") assert response.status_code == 200 diff --git a/tests/unit/api/routes/test_jira_webhook.py b/tests/unit/api/routes/test_jira_webhook.py index cf256b68..99bb3a42 100644 --- a/tests/unit/api/routes/test_jira_webhook.py +++ b/tests/unit/api/routes/test_jira_webhook.py @@ -259,9 +259,7 @@ async def test_standard_task_with_parent_routed_to_parent(self) -> None: @pytest.mark.asyncio @pytest.mark.parametrize("issue_type", ["Task", "Epic"]) - async def test_managed_standalone_issue_bypasses_parent_check( - self, issue_type: str - ) -> None: + async def test_managed_standalone_issue_bypasses_parent_check(self, issue_type: str) -> None: """Managed standalone Task/Epic issues bypass parent checks and queue under their own key.""" webhook = make_jira_webhook(issue_type=issue_type, labels=["forge:managed"]) payload = json.dumps(webhook).encode() @@ -341,6 +339,7 @@ async def test_task_with_managed_label_in_changelog_bypasses_parent_check(self) called_kwargs = mock_producer.publish_once.call_args.kwargs assert called_kwargs["ticket_key"] == "TEST-123" + class TestJiraWebhookParsing: """Tests for Jira webhook payload parsing.""" diff --git a/tests/unit/integrations/agents/test_agent.py b/tests/unit/integrations/agents/test_agent.py index 30c9a841..3602c8dc 100644 --- a/tests/unit/integrations/agents/test_agent.py +++ b/tests/unit/integrations/agents/test_agent.py @@ -1,12 +1,43 @@ """Unit tests for ForgeAgent.""" +import json +from typing import Any from unittest.mock import ANY, AsyncMock, MagicMock, patch import pytest +from langchain_core.callbacks import CallbackManagerForLLMRun +from langchain_core.language_models.chat_models import SimpleChatModel +from langchain_core.messages import BaseMessage from forge.integrations.agents.agent import ForgeAgent +class MockChatModel(SimpleChatModel): + response: str + + def _call( + self, + _messages: list[BaseMessage], + _stop: list[str] | None = None, + _run_manager: CallbackManagerForLLMRun | None = None, + **_kwargs: Any, + ) -> str: + return self.response + + async def _acall( + self, + _messages: list[BaseMessage], + _stop: list[str] | None = None, + _run_manager: CallbackManagerForLLMRun | None = None, + **_kwargs: Any, + ) -> str: + return self.response + + @property + def _llm_type(self) -> str: + return "mock" + + def _model_agent(backend: str, model: str) -> ForgeAgent: agent = ForgeAgent.__new__(ForgeAgent) agent.settings = MagicMock( @@ -212,3 +243,110 @@ def test_get_skill_paths_returns_default_without_ticket_key(): mock_resolver.assert_called_once_with("", ANY) assert result == ["skills/default/"] + + +@pytest.mark.asyncio +async def test_revise_draft_with_feedback_success(): + """Verify that revise_draft_with_feedback properly renders prompt and parses valid JSON.""" + agent = ForgeAgent() + + mock_model = MockChatModel( + response='{"parent_key": "PROJ-1", "items": [{"id": 1, "summary": "Task 1"}]}' + ) + + with patch.object(agent, "_create_model", return_value=mock_model): + result = await agent.revise_draft_with_feedback( + draft_content='{"items": []}', feedback="Add Task 1", context={"ticket_key": "PROJ-1"} + ) + + assert json.loads(result) == {"parent_key": "PROJ-1", "items": [{"id": 1, "summary": "Task 1"}]} + await agent.close() + + +@pytest.mark.asyncio +async def test_revise_draft_with_feedback_markdown_stripping(): + """Verify that revise_draft_with_feedback strips markdown block and preamble.""" + agent = ForgeAgent() + + llm_response = """ + Certainly! Here is the updated JSON: + ```json + { + "items": [ + {"id": 1, "summary": "Task 1"} + ] + } + ``` + Hope this helps! + """ + mock_model = MockChatModel(response=llm_response) + + with patch.object(agent, "_create_model", return_value=mock_model): + result = await agent.revise_draft_with_feedback( + draft_content='{"items": []}', feedback="Add Task 1", context={"ticket_key": "PROJ-1"} + ) + + assert json.loads(result) == {"items": [{"id": 1, "summary": "Task 1"}]} + await agent.close() + + +@pytest.mark.asyncio +async def test_revise_draft_with_feedback_preamble_no_codeblock(): + """Verify that revise_draft_with_feedback strips preamble and postamble without markdown code block.""" + agent = ForgeAgent() + + llm_response = ( + 'The corrected draft is: {"items": [{"id": 1, "summary": "Task 1"}]} please review.' + ) + mock_model = MockChatModel(response=llm_response) + + with patch.object(agent, "_create_model", return_value=mock_model): + result = await agent.revise_draft_with_feedback( + draft_content='{"items": []}', feedback="Add Task 1", context={"ticket_key": "PROJ-1"} + ) + + assert json.loads(result) == {"items": [{"id": 1, "summary": "Task 1"}]} + await agent.close() + + +@pytest.mark.asyncio +async def test_revise_draft_with_feedback_invalid_json(): + """Verify that revise_draft_with_feedback raises ValueError on invalid JSON output.""" + agent = ForgeAgent() + + mock_model = MockChatModel(response="This is not JSON at all.") + + with ( + patch.object(agent, "_create_model", return_value=mock_model), + pytest.raises(ValueError, match="Failed to parse LLM response as valid JSON"), + ): + await agent.revise_draft_with_feedback( + draft_content='{"items": []}', feedback="Add Task 1", context={"ticket_key": "PROJ-1"} + ) + + await agent.close() + + +@pytest.mark.asyncio +async def test_revise_draft_with_feedback_prompt_formatting(): + """Verify that revise_draft_with_feedback properly renders the prompt with input variables.""" + agent = ForgeAgent() + mock_model = MockChatModel(response='{"items": []}') + + with ( + patch( + "forge.integrations.agents.agent.load_prompt", return_value="FORMATTED PROMPT" + ) as mock_load_prompt, + patch.object(agent, "_create_model", return_value=mock_model), + ): + await agent.revise_draft_with_feedback( + draft_content='{"some": "json"}', feedback="Do this", context={"ticket_key": "PROJ-123"} + ) + + mock_load_prompt.assert_called_once_with( + "revision-draft", + draft_content='{"some": "json"}', + feedback="Do this", + context=json.dumps({"ticket_key": "PROJ-123"}, indent=2), + ) + await agent.close() diff --git a/tests/unit/integrations/agents/test_response_parsing.py b/tests/unit/integrations/agents/test_response_parsing.py index e148e5a6..50d7c343 100644 --- a/tests/unit/integrations/agents/test_response_parsing.py +++ b/tests/unit/integrations/agents/test_response_parsing.py @@ -4,7 +4,6 @@ They use realistic AI output samples to test extraction and parsing logic. """ - from forge.integrations.agents.agent import ForgeAgent @@ -322,12 +321,7 @@ def test_expand_nested_dict(self, monkeypatch): monkeypatch.setenv("API_TOKEN", "token123") config = { - "server": { - "url": "${BASE_URL}/v1", - "headers": { - "Authorization": "Bearer ${API_TOKEN}" - } - } + "server": {"url": "${BASE_URL}/v1", "headers": {"Authorization": "Bearer ${API_TOKEN}"}} } result = agent._expand_env_vars(config) diff --git a/tests/unit/integrations/github/test_content_api.py b/tests/unit/integrations/github/test_content_api.py index a7b4fa05..20b00f0a 100644 --- a/tests/unit/integrations/github/test_content_api.py +++ b/tests/unit/integrations/github/test_content_api.py @@ -171,9 +171,7 @@ async def test_returns_none_on_404(self, github_client): response = MagicMock() response.status_code = 404 response.raise_for_status = MagicMock( - side_effect=httpx.HTTPStatusError( - "Not Found", request=MagicMock(), response=response - ) + side_effect=httpx.HTTPStatusError("Not Found", request=MagicMock(), response=response) ) mock_client.get = AsyncMock(return_value=response) diff --git a/tests/unit/integrations/jira/test_client_attachments.py b/tests/unit/integrations/jira/test_client_attachments.py new file mode 100644 index 00000000..6a77b32f --- /dev/null +++ b/tests/unit/integrations/jira/test_client_attachments.py @@ -0,0 +1,147 @@ +"""Unit tests for JiraClient attachment operations.""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +from forge.integrations.jira.client import JiraClient + + +class TestJiraClientAttachments: + """Tests for attachment helper methods in JiraClient.""" + + @pytest.fixture + def mock_client(self): + """Create JiraClient with mocked settings.""" + with patch("forge.integrations.jira.client.get_settings") as mock_settings: + mock_settings.return_value.jira_base_url = "https://test.atlassian.net" + mock_settings.return_value.jira_api_token = MagicMock() + mock_settings.return_value.jira_api_token.get_secret_value.return_value = "token" + mock_settings.return_value.jira_user_email = "test@example.com" + + client = JiraClient() + return client + + @pytest.mark.asyncio + async def test_get_attachments_success(self, mock_client): + """get_attachments successfully fetches and parses attachment list.""" + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "fields": { + "attachment": [ + { + "id": "10001", + "filename": "spec.md", + "content": "https://test.atlassian.net/rest/api/3/attachment/content/10001", + "size": 1234, + }, + { + "id": "10002", + "filename": "design.png", + "content": "https://test.atlassian.net/rest/api/3/attachment/content/10002", + "size": 5678, + }, + ] + } + } + mock_response.raise_for_status = MagicMock() + + with patch.object(mock_client, "_get_client") as mock_get_client: + mock_http = AsyncMock() + mock_http.request = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_http + + attachments = await mock_client.get_attachments("TEST-123") + + assert len(attachments) == 2 + assert attachments[0]["id"] == "10001" + assert attachments[0]["filename"] == "spec.md" + assert ( + attachments[0]["content_url"] + == "https://test.atlassian.net/rest/api/3/attachment/content/10001" + ) + + assert attachments[1]["id"] == "10002" + assert attachments[1]["filename"] == "design.png" + + mock_http.request.assert_called_once_with( + "GET", + "/issue/TEST-123", + params={"fields": "attachment"}, + ) + + @pytest.mark.asyncio + async def test_download_attachment_success(self, mock_client): + """download_attachment successfully fetches binary content.""" + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.content = b"fake binary file content" + mock_response.raise_for_status = MagicMock() + + content_url = "https://test.atlassian.net/rest/api/3/attachment/content/10001" + + with patch.object(mock_client, "_get_client") as mock_get_client: + mock_http = AsyncMock() + mock_http.request = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_http + + content = await mock_client.download_attachment(content_url) + + assert content == b"fake binary file content" + mock_http.request.assert_called_once_with("GET", content_url) + + @pytest.mark.asyncio + async def test_delete_attachment_success(self, mock_client): + """delete_attachment successfully deletes specified attachment.""" + mock_response = MagicMock() + mock_response.status_code = 204 + mock_response.raise_for_status = MagicMock() + + with patch.object(mock_client, "_get_client") as mock_get_client: + mock_http = AsyncMock() + mock_http.request = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_http + + await mock_client.delete_attachment("10001") + + mock_http.request.assert_called_once_with("DELETE", "/attachment/10001") + + @pytest.mark.asyncio + async def test_add_attachment_success(self, mock_client): + """add_attachment successfully uploads file as multipart/form-data and sets token header.""" + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = [ + { + "id": "10003", + "filename": "test-file.json", + "content": "https://test.atlassian.net/rest/api/3/attachment/content/10003", + } + ] + mock_response.raise_for_status = MagicMock() + + file_content = b'{"key": "value"}' + + with patch.object(mock_client, "_get_client") as mock_get_client: + mock_http = AsyncMock() + mock_http.headers = httpx.Headers() + mock_http.request = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_http + + result = await mock_client.add_attachment( + issue_key="TEST-123", + filename="test-file.json", + content=file_content, + ) + + assert result["id"] == "10003" + assert result["filename"] == "test-file.json" + + mock_http.request.assert_called_once() + args, kwargs = mock_http.request.call_args + assert args[0] == "POST" + assert args[1] == "/issue/TEST-123/attachments" + assert kwargs["headers"]["X-Atlassian-Token"] == "no-check" + assert kwargs["files"] == {"file": ("test-file.json", file_content, "application/json")} diff --git a/tests/unit/integrations/langfuse/test_tracing.py b/tests/unit/integrations/langfuse/test_tracing.py index 7f097d7c..88203ea4 100644 --- a/tests/unit/integrations/langfuse/test_tracing.py +++ b/tests/unit/integrations/langfuse/test_tracing.py @@ -7,8 +7,6 @@ from typing import Any from unittest.mock import MagicMock, patch -import pytest - from forge.integrations.langfuse.tracing import ( AsyncLangfuseContext, get_langfuse_config, diff --git a/tests/unit/models/test_bug_state.py b/tests/unit/models/test_bug_state.py index 63f76133..ca732f02 100644 --- a/tests/unit/models/test_bug_state.py +++ b/tests/unit/models/test_bug_state.py @@ -110,7 +110,11 @@ def test_new_fields_serialize_to_json(self): state["rca_options"] = [{"title": "Fix A", "description": "desc", "tradeoffs": "none"}] state["reproducibility_assessment"] = "Unit test feasible" state["selected_fix_option"] = 1 - state["selected_fix_approach"] = {"title": "Fix A", "description": "desc", "tradeoffs": "none"} + state["selected_fix_approach"] = { + "title": "Fix A", + "description": "desc", + "tradeoffs": "none", + } state["plan_content"] = "## Plan\nChange src/auth.py" state["linked_task_keys"] = ["BUG-2", "BUG-3"] state["local_review_verdict"] = "adequate" diff --git a/tests/unit/models/test_draft.py b/tests/unit/models/test_draft.py new file mode 100644 index 00000000..b48faea8 --- /dev/null +++ b/tests/unit/models/test_draft.py @@ -0,0 +1,264 @@ +"""Unit tests for decomposing draft models.""" + +from datetime import UTC, datetime + +import pytest +from pydantic import ValidationError + +from forge.models.draft import DraftItem, ForgeDecompositionDraft + + +class TestDraftItem: + """Tests for DraftItem model validation and serialization.""" + + def test_valid_draft_item(self): + """Verify that a valid DraftItem is successfully created.""" + item = DraftItem( + id=1, + summary="Implement auth route", + description="Create endpoints for signing in and signing up", + repo="auth-service", + acceptance_criteria=[ + "POST /login returns JWT on success", + "POST /register creates a new user", + ], + ) + assert item.id == 1 + assert item.summary == "Implement auth route" + assert item.repo == "auth-service" + assert len(item.acceptance_criteria) == 2 + + def test_invalid_draft_item_types(self): + """Verify that invalid types for DraftItem fields raise ValidationError.""" + with pytest.raises(ValidationError): + DraftItem( + id="invalid_id", # Should be int + summary="Implement auth route", + description="Create endpoints for signing in and signing up", + repo="auth-service", + acceptance_criteria=["Criteria"], + ) + + +class TestForgeDecompositionDraft: + """Tests for ForgeDecompositionDraft validation, serialization, and ID rules.""" + + @pytest.fixture + def valid_items(self) -> list[DraftItem]: + """Return a list of valid, sequential DraftItem objects.""" + return [ + DraftItem( + id=1, + summary="Story 1", + description="Description 1", + repo="repo-a", + acceptance_criteria=["Criteria 1"], + ), + DraftItem( + id=2, + summary="Story 2", + description="Description 2", + repo="repo-b", + acceptance_criteria=["Criteria 2"], + ), + DraftItem( + id=3, + summary="Story 3", + description="Description 3", + repo="repo-a", + acceptance_criteria=["Criteria 3"], + ), + ] + + def test_valid_draft_creation(self, valid_items): + """Verify that a valid draft with unique, sequential IDs can be created.""" + now = datetime.now(UTC) + draft = ForgeDecompositionDraft( + parent_key="PROJ-123", + phase="stories", + items=valid_items, + version=1, + created_at=now, + updated_at=now, + ) + assert draft.parent_key == "PROJ-123" + assert draft.phase == "stories" + assert len(draft.items) == 3 + assert draft.version == 1 + assert draft.created_at == now + assert draft.updated_at == now + + def test_valid_draft_with_empty_items(self): + """Verify that a draft with an empty list of items is valid.""" + now = datetime.now(UTC) + draft = ForgeDecompositionDraft( + parent_key="PROJ-123", + phase="tasks", + items=[], + version=2, + created_at=now, + updated_at=now, + ) + assert draft.items == [] + + def test_invalid_phase(self, valid_items): + """Verify that a phase other than 'stories' or 'tasks' raises ValidationError.""" + now = datetime.now(UTC) + with pytest.raises(ValidationError) as exc_info: + ForgeDecompositionDraft( + parent_key="PROJ-123", + phase="epics", # Invalid phase + items=valid_items, + created_at=now, + updated_at=now, + ) + assert "Input should be 'stories' or 'tasks'" in str(exc_info.value) + + def test_duplicate_ids(self): + """Verify that duplicate item IDs raise ValidationError.""" + now = datetime.now(UTC) + items_with_duplicates = [ + DraftItem( + id=1, + summary="Story 1", + description="Description 1", + repo="repo-a", + acceptance_criteria=[], + ), + DraftItem( + id=1, # Duplicate + summary="Story 2", + description="Description 2", + repo="repo-b", + acceptance_criteria=[], + ), + ] + with pytest.raises(ValidationError) as exc_info: + ForgeDecompositionDraft( + parent_key="PROJ-123", + phase="stories", + items=items_with_duplicates, + created_at=now, + updated_at=now, + ) + assert "Draft item IDs must be unique." in str(exc_info.value) + + def test_non_sequential_ids(self): + """Verify that non-sequential item IDs (gaps) raise ValidationError.""" + now = datetime.now(UTC) + items_with_gap = [ + DraftItem( + id=1, + summary="Story 1", + description="D1", + repo="repo-a", + acceptance_criteria=[], + ), + DraftItem( + id=3, # Gap: missing ID 2 + summary="Story 2", + description="D2", + repo="repo-b", + acceptance_criteria=[], + ), + ] + with pytest.raises(ValidationError) as exc_info: + ForgeDecompositionDraft( + parent_key="PROJ-123", + phase="stories", + items=items_with_gap, + created_at=now, + updated_at=now, + ) + assert "Draft item IDs must be sequential starting from 1." in str(exc_info.value) + + def test_sequential_not_starting_from_one(self): + """Verify that IDs that are sequential but do not start from 1 raise ValidationError.""" + now = datetime.now(UTC) + items_not_starting_at_one = [ + DraftItem( + id=2, # Starts at 2 + summary="Story 1", + description="D1", + repo="repo-a", + acceptance_criteria=[], + ), + DraftItem( + id=3, + summary="Story 2", + description="D2", + repo="repo-b", + acceptance_criteria=[], + ), + ] + with pytest.raises(ValidationError) as exc_info: + ForgeDecompositionDraft( + parent_key="PROJ-123", + phase="stories", + items=items_not_starting_at_one, + created_at=now, + updated_at=now, + ) + assert "Draft item IDs must be sequential starting from 1." in str(exc_info.value) + + def test_unordered_but_valid_ids(self): + """Verify that items with IDs that are unique and sequential starting from 1 are valid even if unordered in input.""" + now = datetime.now(UTC) + unordered_items = [ + DraftItem( + id=2, + summary="Story 2", + description="D2", + repo="repo-b", + acceptance_criteria=[], + ), + DraftItem( + id=1, + summary="Story 1", + description="D1", + repo="repo-a", + acceptance_criteria=[], + ), + ] + draft = ForgeDecompositionDraft( + parent_key="PROJ-123", + phase="stories", + items=unordered_items, + created_at=now, + updated_at=now, + ) + assert len(draft.items) == 2 + # Verify the original list order is preserved (or at least valid) + assert draft.items[0].id == 2 + assert draft.items[1].id == 1 + + def test_serialization_and_deserialization(self, valid_items): + """Verify successful JSON serialization and deserialization of the draft model.""" + now = datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC) + draft = ForgeDecompositionDraft( + parent_key="PROJ-123", + phase="stories", + items=valid_items, + version=1, + created_at=now, + updated_at=now, + ) + + # Serialize to JSON + json_data = draft.model_dump_json() + + # Deserialize back to a new model + restored = ForgeDecompositionDraft.model_validate_json(json_data) + + assert restored.parent_key == draft.parent_key + assert restored.phase == draft.phase + assert restored.version == draft.version + assert restored.created_at == draft.created_at + assert restored.updated_at == draft.updated_at + assert len(restored.items) == len(draft.items) + for original, deserialized in zip(draft.items, restored.items, strict=True): + assert original.id == deserialized.id + assert original.summary == deserialized.summary + assert original.description == deserialized.description + assert original.repo == deserialized.repo + assert original.acceptance_criteria == deserialized.acceptance_criteria diff --git a/tests/unit/orchestrator/gates/test_plan_approval.py b/tests/unit/orchestrator/gates/test_plan_approval.py index 348b06fd..1104f5c8 100644 --- a/tests/unit/orchestrator/gates/test_plan_approval.py +++ b/tests/unit/orchestrator/gates/test_plan_approval.py @@ -57,15 +57,17 @@ def plan_pending_state(self): state["is_paused"] = True return state - def test_routes_to_tasks_on_approval(self, plan_pending_state): + @pytest.mark.asyncio + async def test_routes_to_tasks_on_approval(self, plan_pending_state): """Approved Plan routes to task generation when not paused.""" plan_pending_state["is_paused"] = False - result = route_plan_approval(plan_pending_state) + result = await route_plan_approval(plan_pending_state) assert result == "generate_tasks" - def test_routes_to_regenerate_all_on_full_rejection(self, plan_pending_state): + @pytest.mark.asyncio + async def test_routes_to_regenerate_all_on_full_rejection(self, plan_pending_state): """Full plan rejection routes to regenerate all epics.""" plan_pending_state["context"] = { "labels": ["forge:managed", "forge:plan-pending"], @@ -74,11 +76,12 @@ def test_routes_to_regenerate_all_on_full_rejection(self, plan_pending_state): plan_pending_state["feedback_comment"] = "The epic breakdown doesn't make sense." plan_pending_state["revision_requested"] = True - result = route_plan_approval(plan_pending_state) + result = await route_plan_approval(plan_pending_state) assert result == "regenerate_all_epics" - def test_routes_to_update_single_on_epic_rejection(self, plan_pending_state): + @pytest.mark.asyncio + async def test_routes_to_update_single_on_epic_rejection(self, plan_pending_state): """Single epic rejection routes to update that epic.""" plan_pending_state["context"] = { "labels": ["forge:managed", "forge:plan-pending"], @@ -89,17 +92,18 @@ def test_routes_to_update_single_on_epic_rejection(self, plan_pending_state): plan_pending_state["feedback_comment"] = "Epic 2 needs more detail." plan_pending_state["revision_requested"] = True - result = route_plan_approval(plan_pending_state) + result = await route_plan_approval(plan_pending_state) assert result == "update_single_epic" - def test_routes_to_end_when_pending(self, plan_pending_state): + @pytest.mark.asyncio + async def test_routes_to_end_when_pending(self, plan_pending_state): """Pending Plan without feedback routes to END.""" plan_pending_state["context"] = { "labels": ["forge:managed", "forge:plan-pending"], } - result = route_plan_approval(plan_pending_state) + result = await route_plan_approval(plan_pending_state) assert result == END @@ -118,7 +122,8 @@ def state_with_epics(self): state["epic_keys"] = ["TEST-124", "TEST-125", "TEST-126"] return state - def test_full_regen_deletes_all_epics(self, state_with_epics): + @pytest.mark.asyncio + async def test_full_regen_deletes_all_epics(self, state_with_epics): """Full regeneration affects all epics.""" state_with_epics["context"] = { "labels": ["forge:managed", "forge:plan-pending"], @@ -127,11 +132,12 @@ def test_full_regen_deletes_all_epics(self, state_with_epics): state_with_epics["feedback_comment"] = "Start over with a different approach." state_with_epics["revision_requested"] = True - result = route_plan_approval(state_with_epics) + result = await route_plan_approval(state_with_epics) assert result == "regenerate_all_epics" - def test_single_epic_update_preserves_others(self, state_with_epics): + @pytest.mark.asyncio + async def test_single_epic_update_preserves_others(self, state_with_epics): """Single epic update preserves other epics.""" state_with_epics["context"] = { "labels": ["forge:managed", "forge:plan-pending"], @@ -142,14 +148,15 @@ def test_single_epic_update_preserves_others(self, state_with_epics): state_with_epics["feedback_comment"] = "Just fix this one epic." state_with_epics["revision_requested"] = True - result = route_plan_approval(state_with_epics) + result = await route_plan_approval(state_with_epics) assert result == "update_single_epic" # Other epics should remain in state assert "TEST-124" in state_with_epics["epic_keys"] assert "TEST-126" in state_with_epics["epic_keys"] - def test_partial_approval_scenario(self, state_with_epics): + @pytest.mark.asyncio + async def test_partial_approval_scenario(self, state_with_epics): """Some epics approved, one needs revision.""" # This tests the scenario where user approves some epics # but requests changes to one specific epic @@ -163,7 +170,7 @@ def test_partial_approval_scenario(self, state_with_epics): state_with_epics["feedback_comment"] = "Epic 3 scope is too broad." state_with_epics["revision_requested"] = True - result = route_plan_approval(state_with_epics) + result = await route_plan_approval(state_with_epics) assert result == "update_single_epic" @@ -186,41 +193,172 @@ def plan_pending_state(self): state["is_paused"] = False return state - def test_routes_to_answer_question_when_is_question(self, plan_pending_state): + @pytest.mark.asyncio + async def test_routes_to_answer_question_when_is_question(self, plan_pending_state): """Questions route to answer_question node.""" plan_pending_state["is_question"] = True plan_pending_state["feedback_comment"] = "?Why split into two epics?" - result = route_plan_approval(plan_pending_state) + result = await route_plan_approval(plan_pending_state) assert result == "answer_question" - def test_question_takes_priority_over_revision(self, plan_pending_state): + @pytest.mark.asyncio + async def test_question_takes_priority_over_revision(self, plan_pending_state): """Question routing takes priority over revision routing.""" plan_pending_state["is_question"] = True plan_pending_state["revision_requested"] = True plan_pending_state["feedback_comment"] = "?What's the dependency order?" - result = route_plan_approval(plan_pending_state) + result = await route_plan_approval(plan_pending_state) assert result == "answer_question" - def test_routes_to_regenerate_all_when_feedback_not_question(self, plan_pending_state): + @pytest.mark.asyncio + async def test_routes_to_regenerate_all_when_feedback_not_question(self, plan_pending_state): """Normal feedback routes to regenerate all epics.""" plan_pending_state["is_question"] = False plan_pending_state["revision_requested"] = True plan_pending_state["feedback_comment"] = "Rethink the epic breakdown" - result = route_plan_approval(plan_pending_state) + result = await route_plan_approval(plan_pending_state) assert result == "regenerate_all_epics" - def test_question_without_feedback_does_not_route_to_answer(self, plan_pending_state): + @pytest.mark.asyncio + async def test_question_without_feedback_does_not_route_to_answer(self, plan_pending_state): """is_question alone without feedback_comment doesn't route to answer.""" plan_pending_state["is_question"] = True plan_pending_state["feedback_comment"] = "" - result = route_plan_approval(plan_pending_state) + result = await route_plan_approval(plan_pending_state) # Should proceed to generate_tasks since not paused assert result == "generate_tasks" + + +class TestPlanDraftProvisioning: + """Tests for draft-based ticket provisioning in route_plan_approval.""" + + @pytest.fixture + def approved_plan_state(self): + """Approved plan state waiting for ticket creation.""" + state = create_initial_state( + thread_id="test-thread", + ticket_key="TEST-123", + ticket_type=TicketType.FEATURE, + ) + state["is_paused"] = False + state["epic_keys"] = [] + return state + + @pytest.mark.asyncio + async def test_successful_draft_provisioning(self, approved_plan_state): + """Verify successful download, parsing, skipping excluded items, and deletion on success.""" + from datetime import UTC, datetime + from unittest.mock import AsyncMock, patch + + from forge.models.draft import DraftItem, ForgeDecompositionDraft + + draft_item_1 = DraftItem( + id=1, + summary="Epic One", + description="Details of epic 1", + repo="org/repo-1", + acceptance_criteria=[], + excluded=False, + ) + draft_item_2 = DraftItem( + id=2, + summary="Epic Two", + description="Details of epic 2", + repo="org/repo-2", + acceptance_criteria=[], + excluded=True, # Excluded! + ) + draft = ForgeDecompositionDraft( + parent_key="TEST-123", + phase="stories", + items=[draft_item_1, draft_item_2], + version=1, + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + + with ( + patch("forge.integrations.jira.client.JiraClient") as MockJira, + patch("forge.workflow.utils.draft_manager.DraftManager") as MockDraftManager, + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + + mock_issue = AsyncMock() + mock_issue.project_key = "TEST" + mock_jira.get_issue = AsyncMock(return_value=mock_issue) + mock_jira.create_epic = AsyncMock(return_value="EPIC-101") + + MockDraftManager.get_draft_attachment = AsyncMock(return_value=draft) + MockDraftManager.delete_draft_attachment = AsyncMock() + + result = await route_plan_approval(approved_plan_state) + + assert result == "generate_tasks" + assert approved_plan_state["epic_keys"] == ["EPIC-101"] + + # Verify creations and exclusions + mock_jira.create_epic.assert_called_once_with( + project_key="TEST", + summary="Epic One", + description="Details of epic 1", + parent_key="TEST-123", + labels=["forge:managed", "forge:parent:TEST-123", "repo:org/repo-1"], + ) + + # Verify draft deleted + MockDraftManager.delete_draft_attachment.assert_called_once() + + @pytest.mark.asyncio + async def test_retains_draft_on_failure(self, approved_plan_state): + """Verify that draft attachment is not deleted if epic creation fails midway.""" + from datetime import UTC, datetime + from unittest.mock import AsyncMock, patch + + from forge.models.draft import DraftItem, ForgeDecompositionDraft + + draft_item_1 = DraftItem( + id=1, + summary="Epic One", + description="Details of epic 1", + repo="org/repo-1", + acceptance_criteria=[], + excluded=False, + ) + draft = ForgeDecompositionDraft( + parent_key="TEST-123", + phase="stories", + items=[draft_item_1], + version=1, + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + + with ( + patch("forge.integrations.jira.client.JiraClient") as MockJira, + patch("forge.workflow.utils.draft_manager.DraftManager") as MockDraftManager, + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + + mock_issue = AsyncMock() + mock_issue.project_key = "TEST" + mock_jira.get_issue = AsyncMock(return_value=mock_issue) + mock_jira.create_epic = AsyncMock(side_effect=Exception("Jira failure midway!")) + + MockDraftManager.get_draft_attachment = AsyncMock(return_value=draft) + MockDraftManager.delete_draft_attachment = AsyncMock() + + with pytest.raises(Exception, match="Jira failure midway!"): + await route_plan_approval(approved_plan_state) + + # Deletion should not have been called + MockDraftManager.delete_draft_attachment.assert_not_called() diff --git a/tests/unit/orchestrator/gates/test_task_approval.py b/tests/unit/orchestrator/gates/test_task_approval.py index 5e8b8f7c..81693827 100644 --- a/tests/unit/orchestrator/gates/test_task_approval.py +++ b/tests/unit/orchestrator/gates/test_task_approval.py @@ -59,54 +59,60 @@ def task_pending_state(self): state["is_paused"] = True return state - def test_routes_to_task_router_on_approval(self, task_pending_state): + @pytest.mark.asyncio + async def test_routes_to_task_router_on_approval(self, task_pending_state): """Approved Tasks routes to task router when not paused.""" task_pending_state["is_paused"] = False - result = route_task_approval(task_pending_state) + result = await route_task_approval(task_pending_state) assert result == "task_router" - def test_routes_to_regenerate_all_on_feature_rejection(self, task_pending_state): + @pytest.mark.asyncio + async def test_routes_to_regenerate_all_on_feature_rejection(self, task_pending_state): """Full task rejection routes to regenerate all tasks.""" task_pending_state["feedback_comment"] = "The task breakdown is too coarse." task_pending_state["revision_requested"] = True - result = route_task_approval(task_pending_state) + result = await route_task_approval(task_pending_state) assert result == "regenerate_all_tasks" - def test_routes_to_update_single_on_task_rejection(self, task_pending_state): + @pytest.mark.asyncio + async def test_routes_to_update_single_on_task_rejection(self, task_pending_state): """Single task rejection routes to update that task.""" task_pending_state["current_task_key"] = "TEST-131" task_pending_state["feedback_comment"] = "Task 2 needs more detail." task_pending_state["revision_requested"] = True - result = route_task_approval(task_pending_state) + result = await route_task_approval(task_pending_state) assert result == "update_single_task" - def test_routes_to_regenerate_all_on_epic_sourced_rejection(self, task_pending_state): + @pytest.mark.asyncio + async def test_routes_to_regenerate_all_on_epic_sourced_rejection(self, task_pending_state): """Epic-sourced task feedback routes to regenerate_epic_tasks, not all tasks.""" task_pending_state["current_epic_key"] = "TEST-124" task_pending_state["feedback_comment"] = "Revise the tasks for this epic." task_pending_state["revision_requested"] = True - result = route_task_approval(task_pending_state) + result = await route_task_approval(task_pending_state) assert result == "regenerate_epic_tasks" - def test_feature_level_rejection_still_regenerates_all(self, task_pending_state): + @pytest.mark.asyncio + async def test_feature_level_rejection_still_regenerates_all(self, task_pending_state): """Feature-level feedback (no epic key) still routes to regenerate_all_tasks.""" task_pending_state["current_epic_key"] = None task_pending_state["feedback_comment"] = "The whole task breakdown is wrong." task_pending_state["revision_requested"] = True - result = route_task_approval(task_pending_state) + result = await route_task_approval(task_pending_state) assert result == "regenerate_all_tasks" - def test_epic_rejection_with_empty_body_routes_to_regenerate_epic_tasks( + @pytest.mark.asyncio + async def test_epic_rejection_with_empty_body_routes_to_regenerate_epic_tasks( self, task_pending_state ): """Empty-body '!' on an Epic must not fall through to task_router (approval).""" @@ -114,13 +120,14 @@ def test_epic_rejection_with_empty_body_routes_to_regenerate_epic_tasks( task_pending_state["feedback_comment"] = "" task_pending_state["revision_requested"] = True - result = route_task_approval(task_pending_state) + result = await route_task_approval(task_pending_state) assert result == "regenerate_epic_tasks" - def test_routes_to_end_when_pending(self, task_pending_state): + @pytest.mark.asyncio + async def test_routes_to_end_when_pending(self, task_pending_state): """Pending Tasks without feedback routes to END.""" - result = route_task_approval(task_pending_state) + result = await route_task_approval(task_pending_state) assert result == END @@ -144,41 +151,177 @@ def task_pending_state(self): state["is_paused"] = False return state - def test_routes_to_answer_question_when_is_question(self, task_pending_state): + @pytest.mark.asyncio + async def test_routes_to_answer_question_when_is_question(self, task_pending_state): """Questions route to answer_question node.""" task_pending_state["is_question"] = True task_pending_state["feedback_comment"] = "?Why are there two tasks for this?" - result = route_task_approval(task_pending_state) + result = await route_task_approval(task_pending_state) assert result == "answer_question" - def test_question_takes_priority_over_revision(self, task_pending_state): + @pytest.mark.asyncio + async def test_question_takes_priority_over_revision(self, task_pending_state): """Question routing takes priority over revision routing.""" task_pending_state["is_question"] = True task_pending_state["revision_requested"] = True task_pending_state["feedback_comment"] = "?What's the testing strategy?" - result = route_task_approval(task_pending_state) + result = await route_task_approval(task_pending_state) assert result == "answer_question" - def test_routes_to_regenerate_when_feedback_not_question(self, task_pending_state): + @pytest.mark.asyncio + async def test_routes_to_regenerate_when_feedback_not_question(self, task_pending_state): """Normal feedback routes to regenerate all tasks.""" task_pending_state["is_question"] = False task_pending_state["revision_requested"] = True task_pending_state["feedback_comment"] = "Add more tasks for testing" - result = route_task_approval(task_pending_state) + result = await route_task_approval(task_pending_state) assert result == "regenerate_all_tasks" - def test_question_without_feedback_does_not_route_to_answer(self, task_pending_state): + @pytest.mark.asyncio + async def test_question_without_feedback_does_not_route_to_answer(self, task_pending_state): """is_question alone without feedback_comment doesn't route to answer.""" task_pending_state["is_question"] = True task_pending_state["feedback_comment"] = "" - result = route_task_approval(task_pending_state) + result = await route_task_approval(task_pending_state) # Should proceed to task_router since not paused assert result == "task_router" + + +class TestTaskDraftProvisioning: + """Tests for draft-based ticket provisioning in route_task_approval.""" + + @pytest.fixture + def approved_task_state(self): + """Approved task state waiting for ticket creation.""" + state = create_initial_state( + thread_id="test-thread", + ticket_key="TEST-123", + ticket_type=TicketType.FEATURE, + ) + state["is_paused"] = False + state["epic_keys"] = ["EPIC-124"] + state["task_keys"] = [] + return state + + @pytest.mark.asyncio + async def test_successful_draft_provisioning(self, approved_task_state): + """Verify successful download, parsing, skipping excluded tasks, and deletion on success.""" + from datetime import UTC, datetime + from unittest.mock import AsyncMock, patch + + from forge.models.draft import DraftItem, ForgeDecompositionDraft + + draft_item_1 = DraftItem( + id=1, + summary="Task One", + description="Details of task 1", + repo="org/repo-1", + acceptance_criteria=[], + excluded=False, + epic_key="EPIC-124", + ) + draft_item_2 = DraftItem( + id=2, + summary="Task Two", + description="Details of task 2", + repo="org/repo-2", + acceptance_criteria=[], + excluded=True, # Excluded! + epic_key="EPIC-124", + ) + draft = ForgeDecompositionDraft( + parent_key="TEST-123", + phase="tasks", + items=[draft_item_1, draft_item_2], + version=1, + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + + with ( + patch("forge.integrations.jira.client.JiraClient") as MockJira, + patch("forge.workflow.utils.draft_manager.DraftManager") as MockDraftManager, + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + + mock_issue = AsyncMock() + mock_issue.project_key = "TEST" + mock_jira.get_issue = AsyncMock(return_value=mock_issue) + mock_jira.create_task = AsyncMock(return_value="TASK-201") + + MockDraftManager.get_draft_attachment = AsyncMock(return_value=draft) + MockDraftManager.delete_draft_attachment = AsyncMock() + + result = await route_task_approval(approved_task_state) + + assert result == "task_router" + assert approved_task_state["task_keys"] == ["TASK-201"] + assert approved_task_state["tasks_by_repo"] == {"org/repo-1": ["TASK-201"]} + + # Verify creations and exclusions + mock_jira.create_task.assert_called_once_with( + project_key="TEST", + summary="Task One", + description="Details of task 1", + parent_key="EPIC-124", + labels=["forge:managed", "forge:parent:TEST-123", "repo:org/repo-1"], + ) + + # Verify draft deleted + MockDraftManager.delete_draft_attachment.assert_called_once() + + @pytest.mark.asyncio + async def test_retains_draft_on_failure(self, approved_task_state): + """Verify that draft attachment is not deleted if task creation fails midway.""" + from datetime import UTC, datetime + from unittest.mock import AsyncMock, patch + + from forge.models.draft import DraftItem, ForgeDecompositionDraft + + draft_item_1 = DraftItem( + id=1, + summary="Task One", + description="Details of task 1", + repo="org/repo-1", + acceptance_criteria=[], + excluded=False, + epic_key="EPIC-124", + ) + draft = ForgeDecompositionDraft( + parent_key="TEST-123", + phase="tasks", + items=[draft_item_1], + version=1, + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + + with ( + patch("forge.integrations.jira.client.JiraClient") as MockJira, + patch("forge.workflow.utils.draft_manager.DraftManager") as MockDraftManager, + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + + mock_issue = AsyncMock() + mock_issue.project_key = "TEST" + mock_jira.get_issue = AsyncMock(return_value=mock_issue) + mock_jira.create_task = AsyncMock(side_effect=Exception("Jira task failure!")) + + MockDraftManager.get_draft_attachment = AsyncMock(return_value=draft) + MockDraftManager.delete_draft_attachment = AsyncMock() + + with pytest.raises(Exception, match="Jira task failure!"): + await route_task_approval(approved_task_state) + + # Deletion should not have been called + MockDraftManager.delete_draft_attachment.assert_not_called() diff --git a/tests/unit/orchestrator/gates/test_task_plan_approval.py b/tests/unit/orchestrator/gates/test_task_plan_approval.py index 33df2f08..aa3dce37 100644 --- a/tests/unit/orchestrator/gates/test_task_plan_approval.py +++ b/tests/unit/orchestrator/gates/test_task_plan_approval.py @@ -3,7 +3,6 @@ import pytest from langgraph.graph import END -from forge.models.workflow import TicketType from forge.workflow.gates.task_plan_approval import ( route_task_plan_approval, task_plan_approval_gate, diff --git a/tests/unit/orchestrator/nodes/test_generate_prd.py b/tests/unit/orchestrator/nodes/test_generate_prd.py index a78a1150..0b2b9a61 100644 --- a/tests/unit/orchestrator/nodes/test_generate_prd.py +++ b/tests/unit/orchestrator/nodes/test_generate_prd.py @@ -51,9 +51,7 @@ def mock_jira(self): def mock_agent(self): """Mock ForgeAgent.""" mock = MagicMock() - mock.generate_prd = AsyncMock( - return_value="# PRD\n\n## Overview\nGenerated PRD content." - ) + mock.generate_prd = AsyncMock(return_value="# PRD\n\n## Overview\nGenerated PRD content.") mock.close = AsyncMock() return mock @@ -184,7 +182,9 @@ async def test_regenerates_with_feedback(self, state_with_feedback, mock_jira, m assert "user persona" in call_args.kwargs["feedback"].lower() @pytest.mark.asyncio - async def test_clears_feedback_after_regeneration(self, state_with_feedback, mock_jira, mock_agent): + async def test_clears_feedback_after_regeneration( + self, state_with_feedback, mock_jira, mock_agent + ): """Feedback is cleared after regeneration.""" with ( patch("forge.workflow.nodes.prd_generation.JiraClient", return_value=mock_jira), @@ -221,14 +221,18 @@ async def test_counts_completed_automated_revision( assert result["automated_review_revision_pending"] is False @pytest.mark.asyncio - async def test_stores_in_comment_when_configured(self, state_with_feedback, mock_jira, mock_agent): + async def test_stores_in_comment_when_configured( + self, state_with_feedback, mock_jira, mock_agent + ): """Regenerated PRD is stored as structured comment when jira_store_in_comments is true.""" mock_settings = MagicMock() mock_settings.jira_store_in_comments = True with patch("forge.workflow.nodes.prd_generation.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.prd_generation.ForgeAgent", return_value=mock_agent): - with patch("forge.workflow.nodes.prd_generation.get_settings", return_value=mock_settings): + with patch( + "forge.workflow.nodes.prd_generation.get_settings", return_value=mock_settings + ): await regenerate_prd_with_feedback(state_with_feedback) mock_jira.add_structured_comment.assert_called_once_with( @@ -240,14 +244,18 @@ async def test_stores_in_comment_when_configured(self, state_with_feedback, mock mock_jira.update_description.assert_not_called() @pytest.mark.asyncio - async def test_stores_in_description_when_configured(self, state_with_feedback, mock_jira, mock_agent): + async def test_stores_in_description_when_configured( + self, state_with_feedback, mock_jira, mock_agent + ): """Regenerated PRD updates description when jira_store_in_comments is false.""" mock_settings = MagicMock() mock_settings.jira_store_in_comments = False with patch("forge.workflow.nodes.prd_generation.JiraClient", return_value=mock_jira): with patch("forge.workflow.nodes.prd_generation.ForgeAgent", return_value=mock_agent): - with patch("forge.workflow.nodes.prd_generation.get_settings", return_value=mock_settings): + with patch( + "forge.workflow.nodes.prd_generation.get_settings", return_value=mock_settings + ): await regenerate_prd_with_feedback(state_with_feedback) mock_jira.update_description.assert_called_once_with( diff --git a/tests/unit/orchestrator/test_state.py b/tests/unit/orchestrator/test_state.py index dac398d7..96b09047 100644 --- a/tests/unit/orchestrator/test_state.py +++ b/tests/unit/orchestrator/test_state.py @@ -1,6 +1,5 @@ """Unit tests for workflow state management.""" - from forge.models.workflow import TicketType from forge.workflow.bug.state import create_initial_bug_state from forge.workflow.feature.state import create_initial_feature_state as create_initial_state diff --git a/tests/unit/orchestrator/test_worker.py b/tests/unit/orchestrator/test_worker.py index fbee48e7..383986c6 100644 --- a/tests/unit/orchestrator/test_worker.py +++ b/tests/unit/orchestrator/test_worker.py @@ -1,10 +1,12 @@ """Unit tests for the orchestrator worker.""" +from datetime import UTC from pathlib import Path -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import ANY, AsyncMock, MagicMock, patch import pytest +from forge.models.draft import ForgeDecompositionDraft from forge.models.events import EventSource from forge.orchestrator.worker import ( OrchestratorWorker, @@ -2037,3 +2039,486 @@ async def test_review_response_gate_resume_state_routes_to_implement_review( result = await worker._handle_resume_event(message, state) assert route_review_response(result) == "implement_review" + + +class TestWorkerInteractiveCommentCommandsAndRollback: + """Tests for interactive comment commands and state consistency rollback guard (BR-006).""" + + @pytest.fixture + def worker(self) -> OrchestratorWorker: + return OrchestratorWorker(consumer_name="test-worker") + + @pytest.fixture + def base_message(self) -> QueueMessage: + return QueueMessage( + message_id="msg-123", + event_id="evt-123", + source=EventSource.JIRA, + event_type="jira:issue_updated", + ticket_key="TEST-123", + payload={ + "issue": { + "key": "TEST-123", + "fields": { + "issuetype": {"name": "Feature"}, + "labels": ["forge:managed"], + }, + }, + "comment": { + "id": "10001", + "body": "/forge remove 2", + }, + }, + ) + + @pytest.fixture + def base_state(self) -> dict: + return { + "ticket_key": "TEST-123", + "ticket_type": "Feature", + "current_node": "plan_approval_gate", + "is_paused": True, + "context": {}, + } + + @pytest.fixture + def mock_draft(self) -> ForgeDecompositionDraft: + from datetime import datetime + + from forge.models.draft import DraftItem, ForgeDecompositionDraft + + return ForgeDecompositionDraft( + parent_key="TEST-123", + phase="stories", + items=[ + DraftItem( + id=1, summary="Task 1", description="D1", repo="owner/r", acceptance_criteria=[] + ), + DraftItem( + id=2, summary="Task 2", description="D2", repo="owner/r", acceptance_criteria=[] + ), + ], + version=1, + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + + @pytest.mark.asyncio + async def test_forge_mutation_command_success( + self, + worker: OrchestratorWorker, + base_message: QueueMessage, + base_state: dict, + mock_draft: ForgeDecompositionDraft, + ): + """Worker parses /forge remove 2, mutates draft, saves draft, edits the comment, and stays paused.""" + mock_jira = AsyncMock() + mock_jira.get_attachments.return_value = [{"filename": "forge-stories-draft.json"}] + + mock_review_comment = MagicMock() + mock_review_comment.id = "original_review_comment_id" + mock_review_comment.body = "### 📋 Proposed Epics Draft\nSome details here..." + mock_jira.get_comments.return_value = [mock_review_comment] + + with ( + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + patch( + "forge.workflow.utils.draft_manager.DraftManager.get_draft_attachment", + return_value=mock_draft, + ) as mock_get, + patch( + "forge.workflow.utils.draft_manager.DraftManager.save_draft_attachment", + new_callable=AsyncMock, + ) as mock_save, + ): + result = await worker._handle_resume_event(base_message, base_state) + + # verify get_draft_attachment is called with correct filename + mock_get.assert_called_once_with(mock_jira, "TEST-123", "forge-stories-draft.json") + + # verify save_draft_attachment is called with updated draft where item 2 is removed + mock_save.assert_called_once() + saved_draft = mock_save.call_args[0][2] + assert len(saved_draft.items) == 1 + assert saved_draft.items[0].id == 1 + + # verify edit_comment was called for both original review comment and the command comment + assert mock_jira.edit_comment.call_count == 2 + mock_jira.edit_comment.assert_any_call("TEST-123", "10001", "✅ /forge remove 2") + mock_jira.edit_comment.assert_any_call("TEST-123", "original_review_comment_id", ANY) + + # verify state stays paused + assert result == base_state + assert result["is_paused"] is True + + @pytest.mark.asyncio + async def test_forge_approve_command_success( + self, worker: OrchestratorWorker, base_message: QueueMessage, base_state: dict + ): + """Worker parses /forge approve, sets is_approved=True and workflow label, and unpauses.""" + from forge.models.workflow import ForgeLabel + + base_message.payload["comment"]["body"] = "/forge approve" + mock_jira = AsyncMock() + mock_jira.get_attachments.return_value = [{"filename": "forge-stories-draft.json"}] + + with ( + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + patch( + "forge.orchestrator.worker.provision_epics_from_draft", + new_callable=AsyncMock, + return_value=["EPIC-101", "EPIC-102"], + ) as mock_provision, + ): + result = await worker._handle_resume_event(base_message, base_state) + + # verify workflow label set + mock_jira.set_workflow_label.assert_called_once_with( + "TEST-123", ForgeLabel.PLAN_APPROVED + ) + + # verify provisioning triggered + mock_provision.assert_called_once_with(ANY, mock_jira) + assert result["epic_keys"] == ["EPIC-101", "EPIC-102"] + + # verify state is unpaused + assert result["is_paused"] is False + + @pytest.mark.asyncio + async def test_label_addition_plan_approved_success( + self, worker: OrchestratorWorker, base_message: QueueMessage, base_state: dict + ): + """Label addition event forge:plan-approved triggers Epic provisioning and unpauses.""" + payload = { + **base_message.payload, + "comment": {}, + "changelog": { + "items": [ + { + "field": "labels", + "fromString": "forge:managed forge:plan-pending", + "toString": "forge:managed forge:plan-approved", + } + ] + }, + } + message = QueueMessage( + message_id=base_message.message_id, + event_id=base_message.event_id, + source=base_message.source, + event_type="jira:issue_updated", + ticket_key=base_message.ticket_key, + payload=payload, + ) + + mock_jira = AsyncMock() + with ( + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + patch( + "forge.orchestrator.worker.provision_epics_from_draft", + new_callable=AsyncMock, + return_value=["EPIC-101", "EPIC-102"], + ) as mock_provision, + ): + result = await worker._handle_resume_event(message, base_state) + + # verify provisioning triggered + mock_provision.assert_called_once_with(ANY, mock_jira) + assert result["epic_keys"] == ["EPIC-101", "EPIC-102"] + + # verify state is unpaused + assert result["is_paused"] is False + + @pytest.mark.asyncio + async def test_label_addition_task_approved_success( + self, worker: OrchestratorWorker, base_message: QueueMessage + ): + """Label addition event forge:task-approved triggers Task provisioning and unpauses.""" + task_state = { + "ticket_key": "TEST-123", + "ticket_type": "Feature", + "current_node": "task_approval_gate", + "is_paused": True, + "context": {}, + } + payload = { + **base_message.payload, + "comment": {}, + "changelog": { + "items": [ + { + "field": "labels", + "fromString": "forge:managed forge:task-pending", + "toString": "forge:managed forge:task-approved", + } + ] + }, + } + message = QueueMessage( + message_id=base_message.message_id, + event_id=base_message.event_id, + source=base_message.source, + event_type="jira:issue_updated", + ticket_key=base_message.ticket_key, + payload=payload, + ) + + mock_jira = AsyncMock() + mock_tasks_by_repo = {"org/repo": ["TASK-101", "TASK-102"]} + with ( + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + patch( + "forge.orchestrator.worker.provision_tasks_from_draft", + new_callable=AsyncMock, + return_value=(["TASK-101", "TASK-102"], mock_tasks_by_repo), + ) as mock_provision, + ): + result = await worker._handle_resume_event(message, task_state) + + # verify provisioning triggered + mock_provision.assert_called_once_with(ANY, mock_jira) + assert result["task_keys"] == ["TASK-101", "TASK-102"] + assert result["tasks_by_repo"] == mock_tasks_by_repo + + # verify state is unpaused + assert result["is_paused"] is False + + @pytest.mark.asyncio + async def test_provision_epics_failure_rollback_state_preservation( + self, worker: OrchestratorWorker, base_message: QueueMessage, base_state: dict + ): + """On Epic provisioning failure, keeps state in PENDING_APPROVAL and posts comment.""" + base_message.payload["comment"]["body"] = "/forge approve" + mock_jira = AsyncMock() + mock_jira.get_attachments.return_value = [{"filename": "forge-stories-draft.json"}] + + with ( + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + patch( + "forge.orchestrator.worker.provision_epics_from_draft", + new_callable=AsyncMock, + side_effect=ValueError("Failed to connect to Jira"), + ) as mock_provision, + ): + result = await worker._handle_resume_event(base_message, base_state) + + # verify provisioning was attempted + mock_provision.assert_called_once() + + # verify error comment added + mock_jira.add_comment.assert_called_once_with( + "TEST-123", "❌ Ticket provisioning failed: Failed to connect to Jira" + ) + + # verify state remains paused and unchanged + assert result == base_state + assert result["is_paused"] is True + + @pytest.mark.asyncio + async def test_provision_tasks_failure_rollback_state_preservation( + self, worker: OrchestratorWorker, base_message: QueueMessage + ): + """On Task provisioning failure, keeps state in PENDING_APPROVAL and posts comment.""" + task_state = { + "ticket_key": "TEST-123", + "ticket_type": "Feature", + "current_node": "task_approval_gate", + "is_paused": True, + "context": {}, + } + payload = { + **base_message.payload, + "comment": {}, + "changelog": { + "items": [ + { + "field": "labels", + "fromString": "forge:managed forge:task-pending", + "toString": "forge:managed forge:task-approved", + } + ] + }, + } + message = QueueMessage( + message_id=base_message.message_id, + event_id=base_message.event_id, + source=base_message.source, + event_type="jira:issue_updated", + ticket_key=base_message.ticket_key, + payload=payload, + ) + + mock_jira = AsyncMock() + with ( + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + patch( + "forge.orchestrator.worker.provision_tasks_from_draft", + new_callable=AsyncMock, + side_effect=ValueError("Draft corrupted"), + ) as mock_provision, + ): + result = await worker._handle_resume_event(message, task_state) + + # verify provisioning was attempted + mock_provision.assert_called_once() + + # verify error comment added + mock_jira.add_comment.assert_called_once_with( + "TEST-123", "❌ Ticket provisioning failed: Draft corrupted" + ) + + # verify state remains paused and unchanged + assert result == task_state + assert result["is_paused"] is True + + @pytest.mark.asyncio + async def test_revision_feedback_success( + self, + worker: OrchestratorWorker, + base_message: QueueMessage, + base_state: dict, + mock_draft: ForgeDecompositionDraft, + ): + """Worker parses ! comment, triggers LLM revision chain, saves updated draft, edits comment, and stays paused.""" + + base_message.payload["comment"]["body"] = "! simplify task descriptions" + + mock_jira = AsyncMock() + mock_jira.get_attachments.return_value = [{"filename": "forge-stories-draft.json"}] + + mock_review_comment = MagicMock() + mock_review_comment.id = "original_review_comment_id" + mock_review_comment.body = "### 📋 Proposed Epics Draft\nSome details here..." + mock_jira.get_comments.return_value = [mock_review_comment] + + revised_draft_str = mock_draft.model_copy( + update={"items": [mock_draft.items[0]]} + ).model_dump_json() + + mock_agent = AsyncMock() + mock_agent.revise_draft_with_feedback.return_value = revised_draft_str + mock_agent.close = AsyncMock() + + with ( + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + patch( + "forge.workflow.utils.draft_manager.DraftManager.get_draft_attachment", + return_value=mock_draft, + ), + patch( + "forge.workflow.utils.draft_manager.DraftManager.save_draft_attachment", + new_callable=AsyncMock, + ) as mock_save, + patch("forge.orchestrator.worker.ForgeAgent", return_value=mock_agent), + ): + result = await worker._handle_resume_event(base_message, base_state) + + # verify agent called + mock_agent.revise_draft_with_feedback.assert_called_once() + assert ( + mock_agent.revise_draft_with_feedback.call_args[1]["feedback"] + == "simplify task descriptions" + ) + + # verify save_draft_attachment called with revised draft + mock_save.assert_called_once() + saved_draft = mock_save.call_args[0][2] + assert len(saved_draft.items) == 1 + + # verify edit_comment called + assert mock_jira.edit_comment.call_count == 2 + mock_jira.edit_comment.assert_any_call( + "TEST-123", "10001", "✅ ! simplify task descriptions" + ) + mock_jira.edit_comment.assert_any_call("TEST-123", "original_review_comment_id", ANY) + + # verify state stays paused + assert result == base_state + + @pytest.mark.asyncio + async def test_state_consistency_guard_br_006_on_mutation_failure( + self, + worker: OrchestratorWorker, + base_message: QueueMessage, + base_state: dict, + mock_draft: ForgeDecompositionDraft, + ): + """On mutation failure, rolls back to original draft, keep PENDING_APPROVAL, and post error reply comment.""" + # Malformed command parameters will fail mutation in DraftManager + base_message.payload["comment"]["body"] = "/forge update 2 invalid_param" + + mock_jira = AsyncMock() + mock_jira.get_attachments.return_value = [{"filename": "forge-stories-draft.json"}] + + with ( + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + patch( + "forge.workflow.utils.draft_manager.DraftManager.get_draft_attachment", + return_value=mock_draft, + ), + patch( + "forge.workflow.utils.draft_manager.DraftManager.save_draft_attachment", + new_callable=AsyncMock, + ) as mock_save, + ): + result = await worker._handle_resume_event(base_message, base_state) + + # verify save_draft_attachment is called to roll back (saving original_draft) + mock_save.assert_called_once_with( + mock_jira, "TEST-123", mock_draft, "forge-stories-draft.json" + ) + + # verify error comment is posted + mock_jira.add_comment.assert_called_once() + error_comment = mock_jira.add_comment.call_args[0][1] + assert "❌ Forge command/revision failed:" in error_comment + + # verify state stays paused in PENDING_APPROVAL + assert result == base_state + assert result["is_paused"] is True + + @pytest.mark.asyncio + async def test_state_consistency_guard_br_006_on_revision_failure( + self, + worker: OrchestratorWorker, + base_message: QueueMessage, + base_state: dict, + mock_draft: ForgeDecompositionDraft, + ): + """On LLM revision failure, rolls back to original draft, keep PENDING_APPROVAL, and post error reply comment.""" + base_message.payload["comment"]["body"] = "! make it simpler" + + mock_jira = AsyncMock() + mock_jira.get_attachments.return_value = [{"filename": "forge-stories-draft.json"}] + + mock_agent = AsyncMock() + mock_agent.revise_draft_with_feedback.side_effect = Exception("LLM connection timed out") + mock_agent.close = AsyncMock() + + with ( + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + patch( + "forge.workflow.utils.draft_manager.DraftManager.get_draft_attachment", + return_value=mock_draft, + ), + patch( + "forge.workflow.utils.draft_manager.DraftManager.save_draft_attachment", + new_callable=AsyncMock, + ) as mock_save, + patch("forge.orchestrator.worker.ForgeAgent", return_value=mock_agent), + ): + result = await worker._handle_resume_event(base_message, base_state) + + # verify save_draft_attachment is called to roll back (saving original_draft) + mock_save.assert_called_once_with( + mock_jira, "TEST-123", mock_draft, "forge-stories-draft.json" + ) + + # verify error comment is posted with details + mock_jira.add_comment.assert_called_once() + error_comment = mock_jira.add_comment.call_args[0][1] + assert "❌ Forge command/revision failed: LLM connection timed out" in error_comment + + # verify state stays paused in PENDING_APPROVAL + assert result == base_state + assert result["is_paused"] is True diff --git a/tests/unit/test_main_security.py b/tests/unit/test_main_security.py index 03f1600d..450cdc9b 100644 --- a/tests/unit/test_main_security.py +++ b/tests/unit/test_main_security.py @@ -11,11 +11,7 @@ def test_allow_credentials_is_false(self): from forge.main import create_app app = create_app() - cors = next( - m - for m in app.user_middleware - if m.cls is CORSMiddleware - ) + cors = next(m for m in app.user_middleware if m.cls is CORSMiddleware) assert cors.kwargs["allow_credentials"] is False diff --git a/tests/unit/utils/test_redaction.py b/tests/unit/utils/test_redaction.py index 199879fc..6c8e76ae 100644 --- a/tests/unit/utils/test_redaction.py +++ b/tests/unit/utils/test_redaction.py @@ -5,9 +5,7 @@ def test_redacts_github_token_in_authenticated_url(): token = "gh" + "p_" + "abcdefghijklmnopqrstuvwxyz123456" - text = ( - f"https://x-access-token:{token}@github.com/org/repo.git" - ) + text = f"https://x-access-token:{token}@github.com/org/repo.git" redacted = redact_secrets(text) diff --git a/tests/unit/workflow/bug/test_graph.py b/tests/unit/workflow/bug/test_graph.py index 2ffe8f3b..b6a7d40b 100644 --- a/tests/unit/workflow/bug/test_graph.py +++ b/tests/unit/workflow/bug/test_graph.py @@ -43,9 +43,7 @@ async def test_answer_question_node_receives_bug_rca_artifact_fields(): ], } - with patch( - "forge.workflow.bug.graph.answer_question", new_callable=AsyncMock - ) as mock_answer: + with patch("forge.workflow.bug.graph.answer_question", new_callable=AsyncMock) as mock_answer: mock_answer.side_effect = lambda received: received await graph.compile().ainvoke(state) diff --git a/tests/unit/workflow/bug/test_workflow.py b/tests/unit/workflow/bug/test_workflow.py index f74e8dfa..a825e03d 100644 --- a/tests/unit/workflow/bug/test_workflow.py +++ b/tests/unit/workflow/bug/test_workflow.py @@ -1,7 +1,5 @@ """Tests for BugWorkflow.""" - - from forge.models.workflow import TicketType from forge.workflow.bug.state import create_initial_bug_state @@ -75,6 +73,7 @@ def test_new_fields_have_correct_defaults(self): def test_old_state_without_new_fields_does_not_crash_route_entry(self): """A state dict missing all new fields can be passed to route_entry without KeyError.""" from forge.workflow.bug.graph import route_entry + minimal_old_state = { "ticket_key": "BUG-OLD", "ticket_type": "bug", @@ -88,6 +87,7 @@ def test_old_state_without_new_fields_does_not_crash_route_entry(self): def test_rca_approval_gate_checkpoint_maps_correctly(self): """In-flight state with current_node='rca_approval_gate' routes to rca_option_gate.""" from forge.workflow.bug.graph import route_entry + state = { "ticket_key": "BUG-OLD", "current_node": "rca_approval_gate", @@ -98,6 +98,7 @@ def test_rca_approval_gate_checkpoint_maps_correctly(self): def test_new_fields_not_required_for_route_entry(self): """route_entry handles state dicts missing new fields — uses .get() throughout.""" from forge.workflow.bug.graph import route_entry + for node, expected in [ ("triage_check", "triage_check"), ("analyze_bug", "analyze_bug"), @@ -114,6 +115,7 @@ class TestTasksByRepoInBugState: def test_tasks_by_repo_declared_in_bug_state_annotations(self): """tasks_by_repo is declared in BugState so LangGraph includes it in the checkpoint schema.""" from forge.workflow.bug.state import BugState + all_annotations: dict = {} for cls in BugState.__mro__: all_annotations.update(getattr(cls, "__annotations__", {})) @@ -134,6 +136,7 @@ class TestNewStateFixtures: def test_state_triage_pending_has_correct_fields(self): """STATE_TRIAGE_PENDING represents a paused triage state correctly.""" from tests.fixtures.workflow_states import STATE_TRIAGE_PENDING + assert STATE_TRIAGE_PENDING["is_paused"] is True assert STATE_TRIAGE_PENDING["current_node"] == "triage_gate" assert STATE_TRIAGE_PENDING["triage_passed"] is False @@ -142,6 +145,7 @@ def test_state_triage_pending_has_correct_fields(self): def test_state_rca_option_pending_has_options(self): """STATE_RCA_OPTION_PENDING has at least 2 RCA options with required keys.""" from tests.fixtures.workflow_states import STATE_RCA_OPTION_PENDING + options = STATE_RCA_OPTION_PENDING.get("rca_options", []) assert len(options) >= 2 for opt in options: @@ -152,19 +156,20 @@ def test_state_rca_option_pending_has_options(self): def test_state_bug_plan_pending_has_plan_content(self): """STATE_BUG_PLAN_PENDING has non-empty plan_content.""" from tests.fixtures.workflow_states import STATE_BUG_PLAN_PENDING + assert STATE_BUG_PLAN_PENDING["current_node"] == "plan_approval_gate" assert STATE_BUG_PLAN_PENDING.get("plan_content", "") def test_triage_pending_fixture_routes_to_triage_gate(self): """STATE_TRIAGE_PENDING route_entry returns 'triage_gate'.""" + from forge.workflow.bug.graph import route_entry from tests.fixtures.workflow_states import STATE_TRIAGE_PENDING - from forge.workflow.bug.graph import route_entry assert route_entry(STATE_TRIAGE_PENDING) == "triage_gate" def test_rca_option_pending_fixture_routes_to_rca_option_gate(self): """STATE_RCA_OPTION_PENDING route_entry returns 'rca_option_gate'.""" + from forge.workflow.bug.graph import route_entry from tests.fixtures.workflow_states import STATE_RCA_OPTION_PENDING - from forge.workflow.bug.graph import route_entry assert route_entry(STATE_RCA_OPTION_PENDING) == "rca_option_gate" diff --git a/tests/unit/workflow/feature/test_prd_pr_state.py b/tests/unit/workflow/feature/test_prd_pr_state.py index 103d2f54..a3dd0d68 100644 --- a/tests/unit/workflow/feature/test_prd_pr_state.py +++ b/tests/unit/workflow/feature/test_prd_pr_state.py @@ -1,7 +1,7 @@ """Tests for PRD PR state fields.""" from forge.models.workflow import TicketType -from forge.workflow.feature.state import FeatureState, create_initial_feature_state +from forge.workflow.feature.state import create_initial_feature_state class TestPrdPrStateFields: diff --git a/tests/unit/workflow/nodes/test_ci_attempt_tracking.py b/tests/unit/workflow/nodes/test_ci_attempt_tracking.py index 9cad5d96..538df5ba 100644 --- a/tests/unit/workflow/nodes/test_ci_attempt_tracking.py +++ b/tests/unit/workflow/nodes/test_ci_attempt_tracking.py @@ -1,12 +1,12 @@ """Unit tests for CI attempt tracking (AISOS-654).""" -import pytest from unittest.mock import AsyncMock, MagicMock, patch +import pytest + 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 ─────────────────────────────────────────────────────────────────── diff --git a/tests/unit/workflow/nodes/test_code_review.py b/tests/unit/workflow/nodes/test_code_review.py index e56ac199..618187e6 100644 --- a/tests/unit/workflow/nodes/test_code_review.py +++ b/tests/unit/workflow/nodes/test_code_review.py @@ -33,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", @@ -61,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", @@ -83,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", @@ -143,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 @@ -190,13 +198,19 @@ async def test_updates_pr_when_description_is_inaccurate(self, state): 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.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"), + ): await sync_pr_description( - state, _git_mock(), - owner="org", repo="repo", pr_number=42, attempt=2, + state, + _git_mock(), + owner="org", + repo="repo", + pr_number=42, + attempt=2, ) github.update_pull_request.assert_called_once_with("org", "repo", 42, body=updated) @@ -215,13 +229,19 @@ async def test_skips_when_body_unchanged(self, state): 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.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"), + ): await sync_pr_description( - state, _git_mock(), - owner="org", repo="repo", pr_number=42, attempt=2, + state, + _git_mock(), + owner="org", + repo="repo", + pr_number=42, + attempt=2, ) github.update_pull_request.assert_not_called() @@ -234,12 +254,18 @@ async def test_skips_when_no_commits(self, state): github, jira = _github_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.GitHubClient", return_value=github), + 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(""), + owner="org", + repo="repo", + pr_number=42, + attempt=1, ) MockAgent.assert_not_called() @@ -251,8 +277,12 @@ async def test_skips_when_no_pr_number(self, state): with patch("forge.workflow.nodes.code_review.GitHubClient") as MockGH: await sync_pr_description( - state, MagicMock(), - owner="org", repo="repo", pr_number=None, attempt=1, + state, + MagicMock(), + owner="org", + repo="repo", + pr_number=None, + attempt=1, ) MockGH.assert_not_called() @@ -268,13 +298,19 @@ async def test_error_does_not_propagate(self, state): 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.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"), + ): await sync_pr_description( - state, _git_mock(), - owner="org", repo="repo", pr_number=42, attempt=1, + state, + _git_mock(), + owner="org", + repo="repo", + pr_number=42, + attempt=1, ) github.update_pull_request.assert_not_called() @@ -290,13 +326,19 @@ async def test_audit_comment_labels_initial_create(self, state): 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.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"), + ): await sync_pr_description( - state, _git_mock(), - owner="org", repo="repo", pr_number=42, attempt=0, + state, + _git_mock(), + owner="org", + repo="repo", + pr_number=42, + attempt=0, ) comment_text = jira.add_comment.call_args[0][1] @@ -343,18 +385,24 @@ 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.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, + ): await create_pull_request(state) mock_sync.assert_called_once() diff --git a/tests/unit/workflow/nodes/test_create_pr_bug.py b/tests/unit/workflow/nodes/test_create_pr_bug.py index 663f7be1..4a0b510c 100644 --- a/tests/unit/workflow/nodes/test_create_pr_bug.py +++ b/tests/unit/workflow/nodes/test_create_pr_bug.py @@ -69,7 +69,9 @@ def test_qualitative_review_failed_adds_warning(self): def test_no_warning_when_review_passed(self): """qualitative_review_failed=False → no warning block.""" - body = _build_pr_body(_bug_state(qualitative_review_failed=False), implemented_tasks=["BUG-50"]) + body = _build_pr_body( + _bug_state(qualitative_review_failed=False), implemented_tasks=["BUG-50"] + ) assert "automated qualitative review" not in body.lower() def test_warning_and_release_note_both_appear_when_review_failed(self): diff --git a/tests/unit/workflow/nodes/test_epic_decomposition.py b/tests/unit/workflow/nodes/test_epic_decomposition.py index 8786542c..ff4772dd 100644 --- a/tests/unit/workflow/nodes/test_epic_decomposition.py +++ b/tests/unit/workflow/nodes/test_epic_decomposition.py @@ -7,6 +7,7 @@ from forge.integrations.jira.client import MissingProjectConfig from forge.models.workflow import ForgeLabel from forge.workflow.nodes.epic_decomposition import decompose_epics, regenerate_all_epics +from forge.workflow.utils.draft_manager import DraftManager @pytest.fixture @@ -17,6 +18,7 @@ def base_state(): "qa_history": [], "generation_context": {}, "retry_count": 0, + "yolo_mode": True, } @@ -115,7 +117,9 @@ async def test_blocks_and_comments_when_forge_repos_missing(self, base_state, mo patch("forge.workflow.nodes.epic_decomposition.JiraClient") as MockJira, patch("forge.workflow.nodes.epic_decomposition.ForgeAgent") as MockAgent, patch("forge.workflow.nodes.epic_decomposition.post_qa_summary_if_needed"), - patch("forge.workflow.nodes.epic_decomposition.get_settings", return_value=mock_settings), + patch( + "forge.workflow.nodes.epic_decomposition.get_settings", return_value=mock_settings + ), ): mock_jira = AsyncMock() MockJira.return_value = mock_jira @@ -136,9 +140,7 @@ async def test_blocks_and_comments_when_forge_repos_missing(self, base_state, mo assert "forge.repos" in comment_text assert "forge:retry" in comment_text - mock_jira.set_workflow_label.assert_called_once_with( - "MYPROJ-1", ForgeLabel.BLOCKED - ) + mock_jira.set_workflow_label.assert_called_once_with("MYPROJ-1", ForgeLabel.BLOCKED) assert result["last_error"] assert result["current_node"] == "decompose_epics" @@ -153,7 +155,9 @@ async def test_blocks_and_comments_when_forge_repos_malformed(self, base_state, patch("forge.workflow.nodes.epic_decomposition.JiraClient") as MockJira, patch("forge.workflow.nodes.epic_decomposition.ForgeAgent") as MockAgent, patch("forge.workflow.nodes.epic_decomposition.post_qa_summary_if_needed"), - patch("forge.workflow.nodes.epic_decomposition.get_settings", return_value=mock_settings), + patch( + "forge.workflow.nodes.epic_decomposition.get_settings", return_value=mock_settings + ), ): mock_jira = AsyncMock() MockJira.return_value = mock_jira @@ -171,9 +175,7 @@ async def test_blocks_and_comments_when_forge_repos_malformed(self, base_state, result = await decompose_epics(base_state) - mock_jira.set_workflow_label.assert_called_once_with( - "MYPROJ-1", ForgeLabel.BLOCKED - ) + mock_jira.set_workflow_label.assert_called_once_with("MYPROJ-1", ForgeLabel.BLOCKED) assert result["last_error"] @@ -255,3 +257,182 @@ async def test_regenerate_all_epics_clears_revision_flags_after_new_epics( assert result["current_node"] == "plan_approval_gate" assert result["revision_requested"] is False assert result["feedback_comment"] is None + + +class TestDecomposeEpicsDraftReview: + """Tests for the non-YOLO draft review gate flow in decompose_epics.""" + + @pytest.mark.asyncio + async def test_decompose_epics_draft_review_flow_success( + self, base_state, mock_issue, mock_epics_data + ): + """When yolo_mode is False, decomposes epics into a draft JSON, deletes old attachments, saves new one, posts comment, and pauses.""" + state = {**base_state, "yolo_mode": False} + + with ( + patch("forge.workflow.nodes.epic_decomposition.JiraClient") as MockJira, + patch("forge.workflow.nodes.epic_decomposition.ForgeAgent") as MockAgent, + patch("forge.workflow.nodes.epic_decomposition.DraftManager") as MockDraftManager, + patch("forge.workflow.nodes.epic_decomposition.post_qa_summary_if_needed"), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(return_value=mock_issue) + mock_jira.get_labels = AsyncMock(return_value=[]) + mock_jira.get_project_repos = AsyncMock(return_value=["acme/backend"]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + mock_agent.generate_epics = AsyncMock(return_value=mock_epics_data) + + MockDraftManager.delete_draft_attachment = AsyncMock() + MockDraftManager.save_draft_attachment = AsyncMock() + MockDraftManager.format_review_comment.side_effect = DraftManager.format_review_comment + + result = await decompose_epics(state) + + # 1. Verify DraftManager deleted any existing forge-stories-draft.json first + MockDraftManager.delete_draft_attachment.assert_called_once_with( + mock_jira, "MYPROJ-1", "forge-stories-draft.json" + ) + + # 2. Verify DraftManager saved the new draft + MockDraftManager.save_draft_attachment.assert_called_once() + saved_draft = MockDraftManager.save_draft_attachment.call_args[0][2] + assert saved_draft.parent_key == "MYPROJ-1" + assert saved_draft.phase == "stories" + assert len(saved_draft.items) == 1 + assert saved_draft.items[0].summary == "Epic One" + assert saved_draft.items[0].description == "Do stuff." + assert saved_draft.items[0].repo == "acme/backend" + + # 3. Verify formatted comment posted + assert mock_jira.add_comment.call_count == 2 + comment_text = mock_jira.add_comment.call_args_list[1][0][1] + assert "### 📋 Proposed Epics Draft" in comment_text + assert "Epic One" in comment_text + assert "acme/backend" in comment_text + + # 4. Verify workflow label updated to PLAN_PENDING + mock_jira.set_workflow_label.assert_called_once_with("MYPROJ-1", ForgeLabel.PLAN_PENDING) + + # 5. Verify state transitions to plan_approval_gate and pauses + assert result["current_node"] == "plan_approval_gate" + assert result["is_paused"] is True + assert result["epic_keys"] == [] + + @pytest.mark.asyncio + async def test_decompose_epics_draft_review_truncation_limits(self, base_state, mock_issue): + """When the item list has > 15 elements, comment falls back to a condensed table.""" + state = {**base_state, "yolo_mode": False} + + # Mock 16 items + many_epics_data = [ + {"summary": f"Epic {i}", "plan": f"Plan {i}", "repo": f"repo-{i}"} for i in range(1, 17) + ] + + with ( + patch("forge.workflow.nodes.epic_decomposition.JiraClient") as MockJira, + patch("forge.workflow.nodes.epic_decomposition.ForgeAgent") as MockAgent, + patch("forge.workflow.nodes.epic_decomposition.DraftManager") as MockDraftManager, + patch("forge.workflow.nodes.epic_decomposition.post_qa_summary_if_needed"), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(return_value=mock_issue) + mock_jira.get_labels = AsyncMock(return_value=[]) + mock_jira.get_project_repos = AsyncMock(return_value=["acme/backend"]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + mock_agent.generate_epics = AsyncMock(return_value=many_epics_data) + + MockDraftManager.delete_draft_attachment = AsyncMock() + MockDraftManager.save_draft_attachment = AsyncMock() + MockDraftManager.format_review_comment.side_effect = DraftManager.format_review_comment + + await decompose_epics(state) + + # Verify comment is in condensed table format + assert mock_jira.add_comment.call_count == 2 + comment_text = mock_jira.add_comment.call_args_list[1][0][1] + assert "### 📋 Proposed Epics Draft (Condensed)" in comment_text + assert "Warning" in comment_text + assert "forge-stories-draft.json" in comment_text + # Condensed table should only show IDs, summaries, and target repos + # Detailed descriptions/plans (like Plan 1) should NOT be in the comment + assert "Plan 1" not in comment_text + assert "Epic 1" in comment_text + assert "repo-1" in comment_text + + @pytest.mark.asyncio + async def test_decompose_epics_draft_review_truncation_characters(self, base_state, mock_issue): + """When comment exceeds 32,767 characters, comment falls back to a condensed table.""" + state = {**base_state, "yolo_mode": False} + + # Mock huge description to exceed character limit + huge_epics_data = [{"summary": "Epic One", "plan": "A" * 35000, "repo": "acme/backend"}] + + with ( + patch("forge.workflow.nodes.epic_decomposition.JiraClient") as MockJira, + patch("forge.workflow.nodes.epic_decomposition.ForgeAgent") as MockAgent, + patch("forge.workflow.nodes.epic_decomposition.DraftManager") as MockDraftManager, + patch("forge.workflow.nodes.epic_decomposition.post_qa_summary_if_needed"), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(return_value=mock_issue) + mock_jira.get_labels = AsyncMock(return_value=[]) + mock_jira.get_project_repos = AsyncMock(return_value=["acme/backend"]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + mock_agent.generate_epics = AsyncMock(return_value=huge_epics_data) + + MockDraftManager.delete_draft_attachment = AsyncMock() + MockDraftManager.save_draft_attachment = AsyncMock() + MockDraftManager.format_review_comment.side_effect = DraftManager.format_review_comment + + await decompose_epics(state) + + # Verify comment is in condensed table format due to length + assert mock_jira.add_comment.call_count == 2 + comment_text = mock_jira.add_comment.call_args_list[1][0][1] + assert "### 📋 Proposed Epics Draft (Condensed)" in comment_text + assert "Warning" in comment_text + assert "forge-stories-draft.json" in comment_text + assert "A" * 35000 not in comment_text + assert "Epic One" in comment_text + assert "acme/backend" in comment_text + + @pytest.mark.asyncio + async def test_decompose_epics_empty_data_retry(self, base_state, mock_issue): + """When empty epics_data is returned, returns a retry state with last_error and retry_count incremented.""" + state = {**base_state} + + with ( + patch("forge.workflow.nodes.epic_decomposition.JiraClient") as MockJira, + patch("forge.workflow.nodes.epic_decomposition.ForgeAgent") as MockAgent, + patch("forge.workflow.nodes.epic_decomposition.post_qa_summary_if_needed"), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(return_value=mock_issue) + mock_jira.get_labels = AsyncMock(return_value=[]) + mock_jira.get_project_repos = AsyncMock(return_value=["acme/backend"]) + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + mock_agent.generate_epics = AsyncMock(return_value=[]) + + result = await decompose_epics(state) + + assert result["current_node"] == "decompose_epics" + assert result["retry_count"] == 1 + assert "Epic generation returned no results" in result["last_error"] diff --git a/tests/unit/workflow/nodes/test_escalate_to_blocked.py b/tests/unit/workflow/nodes/test_escalate_to_blocked.py index 103cfec9..de2954d2 100644 --- a/tests/unit/workflow/nodes/test_escalate_to_blocked.py +++ b/tests/unit/workflow/nodes/test_escalate_to_blocked.py @@ -3,6 +3,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest + from tests.fixtures.workflow_states import make_workflow_state @@ -29,10 +30,12 @@ def mock_jira(): jira = MagicMock() jira.set_workflow_label = AsyncMock() jira.add_comment = AsyncMock() - jira.get_issue = AsyncMock(return_value=MagicMock( - reporter="reporter@example.com", - assignee="assignee@example.com", - )) + jira.get_issue = AsyncMock( + return_value=MagicMock( + reporter="reporter@example.com", + assignee="assignee@example.com", + ) + ) jira.close = AsyncMock() return jira @@ -45,8 +48,10 @@ async def test_sets_is_blocked_true(self, state_at_ci, mock_jira): """Result state has is_blocked=True.""" from forge.workflow.nodes.ci_evaluator import escalate_to_blocked - with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()): + with ( + patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()), + ): result = await escalate_to_blocked(state_at_ci) assert result.get("is_blocked") is True @@ -56,8 +61,10 @@ async def test_sets_is_blocked_from_workspace_failure(self, state_at_workspace, """is_blocked=True regardless of which node triggered escalation.""" from forge.workflow.nodes.ci_evaluator import escalate_to_blocked - with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()): + with ( + patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()), + ): result = await escalate_to_blocked(state_at_workspace) assert result.get("is_blocked") is True @@ -71,8 +78,10 @@ async def test_preserves_current_node_at_ci(self, state_at_ci, mock_jira): """current_node stays 'ci_evaluator' after CI exhaustion escalation.""" from forge.workflow.nodes.ci_evaluator import escalate_to_blocked - with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()): + with ( + patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()), + ): result = await escalate_to_blocked(state_at_ci) assert result["current_node"] == "ci_evaluator" @@ -82,8 +91,10 @@ async def test_preserves_current_node_at_workspace(self, state_at_workspace, moc """current_node stays 'setup_workspace' after workspace failure.""" from forge.workflow.nodes.ci_evaluator import escalate_to_blocked - with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()): + with ( + patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()), + ): result = await escalate_to_blocked(state_at_workspace) assert result["current_node"] == "setup_workspace" @@ -93,8 +104,10 @@ async def test_does_not_set_current_node_to_complete(self, state_at_ci, mock_jir """current_node must never be set to 'complete' by escalation.""" from forge.workflow.nodes.ci_evaluator import escalate_to_blocked - with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()): + with ( + patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()), + ): result = await escalate_to_blocked(state_at_ci) assert result["current_node"] != "complete" @@ -109,8 +122,10 @@ async def test_sets_blocked_jira_label(self, state_at_ci, mock_jira): from forge.models.workflow import ForgeLabel from forge.workflow.nodes.ci_evaluator import escalate_to_blocked - with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()): + with ( + patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()), + ): await escalate_to_blocked(state_at_ci) mock_jira.set_workflow_label.assert_called_once_with( @@ -122,8 +137,10 @@ async def test_sets_ci_status_to_blocked(self, state_at_ci, mock_jira): """ci_status is set to 'blocked' in the returned state.""" from forge.workflow.nodes.ci_evaluator import escalate_to_blocked - with patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()): + with ( + patch("forge.workflow.nodes.ci_evaluator.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.error_handler.notify_error", AsyncMock()), + ): result = await escalate_to_blocked(state_at_ci) assert result.get("ci_status") == "blocked" diff --git a/tests/unit/workflow/nodes/test_generation_context.py b/tests/unit/workflow/nodes/test_generation_context.py index 1c7d2887..32c75c52 100644 --- a/tests/unit/workflow/nodes/test_generation_context.py +++ b/tests/unit/workflow/nodes/test_generation_context.py @@ -54,9 +54,7 @@ async def test_generate_prd_stores_generation_context(self): ) mock_agent = create_mock_forge_agent() - mock_agent.generate_prd = AsyncMock( - return_value="# Generated PRD\n\nContent here." - ) + mock_agent.generate_prd = AsyncMock(return_value="# Generated PRD\n\nContent here.") state = create_initial_feature_state( ticket_key="TEST-123", @@ -103,9 +101,7 @@ async def test_generate_prd_preserves_existing_context(self): ) mock_agent = create_mock_forge_agent() - mock_agent.generate_prd = AsyncMock( - return_value="# PRD Content" - ) + mock_agent.generate_prd = AsyncMock(return_value="# PRD Content") state = create_initial_feature_state( ticket_key="TEST-123", @@ -141,9 +137,7 @@ async def test_generate_spec_stores_generation_context(self): mock_jira = create_mock_jira_client() mock_agent = create_mock_forge_agent() - mock_agent.generate_spec = AsyncMock( - return_value="# Generated Spec\n\nContent here." - ) + mock_agent.generate_spec = AsyncMock(return_value="# Generated Spec\n\nContent here.") state = create_initial_feature_state( ticket_key="TEST-123", @@ -182,9 +176,7 @@ async def test_generate_spec_preserves_prd_context(self): mock_jira = create_mock_jira_client() mock_agent = create_mock_forge_agent() - mock_agent.generate_spec = AsyncMock( - return_value="# Spec Content" - ) + mock_agent.generate_spec = AsyncMock(return_value="# Spec Content") state = create_initial_feature_state( ticket_key="TEST-123", diff --git a/tests/unit/workflow/nodes/test_implementation_status_instrumentation.py b/tests/unit/workflow/nodes/test_implementation_status_instrumentation.py index 482f8470..f76ea8a7 100644 --- a/tests/unit/workflow/nodes/test_implementation_status_instrumentation.py +++ b/tests/unit/workflow/nodes/test_implementation_status_instrumentation.py @@ -5,7 +5,6 @@ correct parameters, independent of the Jira client implementation. """ -from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -70,9 +69,7 @@ async def test_post_status_comment_called_at_start_with_correct_params(self): with ( patch("forge.workflow.nodes.implementation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.implementation.ContainerRunner", return_value=mock_runner), - patch( - "forge.workflow.nodes.implementation.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.implementation.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() result = await implement_task(state) @@ -84,7 +81,9 @@ async def test_post_status_comment_called_at_start_with_correct_params(self): first_call = mock_post_status.call_args_list[0] assert first_call[0][0] == mock_jira # JiraClient instance assert first_call[0][1] == "TASK-1" # task_key - assert first_call[0][2] == "🔨 Forge started implementing [TASK-1]: Task summary" # start message + assert ( + first_call[0][2] == "🔨 Forge started implementing [TASK-1]: Task summary" + ) # start message @pytest.mark.asyncio async def test_post_status_comment_called_before_container_execution(self): @@ -151,9 +150,7 @@ async def test_post_status_comment_called_at_completion_on_success(self): with ( patch("forge.workflow.nodes.implementation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.implementation.ContainerRunner", return_value=mock_runner), - patch( - "forge.workflow.nodes.implementation.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.implementation.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() result = await implement_task(state) @@ -166,8 +163,7 @@ async def test_post_status_comment_called_at_completion_on_success(self): assert second_call[0][0] == mock_jira # JiraClient instance assert second_call[0][1] == "TASK-1" # task_key assert ( - second_call[0][2] - == "✅ Implementation complete. Running local code review before PR." + second_call[0][2] == "✅ Implementation complete. Running local code review before PR." ) @pytest.mark.asyncio @@ -188,9 +184,7 @@ async def test_post_status_comment_not_called_at_completion_on_failure(self): with ( patch("forge.workflow.nodes.implementation.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.implementation.ContainerRunner", return_value=mock_runner), - patch( - "forge.workflow.nodes.implementation.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.implementation.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() result = await implement_task(state) @@ -226,9 +220,7 @@ async def test_multiple_tasks_use_correct_task_key_for_each_comment(self): with ( patch("forge.workflow.nodes.implementation.JiraClient", return_value=mock_jira1), patch("forge.workflow.nodes.implementation.ContainerRunner", return_value=mock_runner1), - patch( - "forge.workflow.nodes.implementation.post_status_comment" - ) as mock_post_status1, + patch("forge.workflow.nodes.implementation.post_status_comment") as mock_post_status1, ): mock_post_status1.return_value = AsyncMock() result1 = await implement_task(state1) @@ -247,9 +239,7 @@ async def test_multiple_tasks_use_correct_task_key_for_each_comment(self): with ( patch("forge.workflow.nodes.implementation.JiraClient", return_value=mock_jira2), patch("forge.workflow.nodes.implementation.ContainerRunner", return_value=mock_runner2), - patch( - "forge.workflow.nodes.implementation.post_status_comment" - ) as mock_post_status2, + patch("forge.workflow.nodes.implementation.post_status_comment") as mock_post_status2, ): mock_post_status2.return_value = AsyncMock() result2 = await implement_task(state2) @@ -268,9 +258,7 @@ async def test_multiple_tasks_use_correct_task_key_for_each_comment(self): with ( patch("forge.workflow.nodes.implementation.JiraClient", return_value=mock_jira3), patch("forge.workflow.nodes.implementation.ContainerRunner", return_value=mock_runner3), - patch( - "forge.workflow.nodes.implementation.post_status_comment" - ) as mock_post_status3, + patch("forge.workflow.nodes.implementation.post_status_comment") as mock_post_status3, ): mock_post_status3.return_value = AsyncMock() result3 = await implement_task(state3) @@ -299,9 +287,7 @@ async def test_multiple_tasks_mixed_success_failure_correct_task_keys(self): with ( patch("forge.workflow.nodes.implementation.JiraClient", return_value=mock_jira1), patch("forge.workflow.nodes.implementation.ContainerRunner", return_value=mock_runner1), - patch( - "forge.workflow.nodes.implementation.post_status_comment" - ) as mock_post_status1, + patch("forge.workflow.nodes.implementation.post_status_comment") as mock_post_status1, ): mock_post_status1.return_value = AsyncMock() result1 = await implement_task(state1) @@ -322,9 +308,7 @@ async def test_multiple_tasks_mixed_success_failure_correct_task_keys(self): with ( patch("forge.workflow.nodes.implementation.JiraClient", return_value=mock_jira2), patch("forge.workflow.nodes.implementation.ContainerRunner", return_value=mock_runner2), - patch( - "forge.workflow.nodes.implementation.post_status_comment" - ) as mock_post_status2, + patch("forge.workflow.nodes.implementation.post_status_comment") as mock_post_status2, ): mock_post_status2.return_value = AsyncMock() result2 = await implement_task(state2) @@ -333,5 +317,6 @@ async def test_multiple_tasks_mixed_success_failure_correct_task_keys(self): assert mock_post_status2.call_count == 1 assert mock_post_status2.call_args_list[0][0][1] == "TASK-2" assert ( - mock_post_status2.call_args_list[0][0][2] == "🔨 Forge started implementing [TASK-2]: Task summary" + mock_post_status2.call_args_list[0][0][2] + == "🔨 Forge started implementing [TASK-2]: Task summary" ) diff --git a/tests/unit/workflow/nodes/test_local_review_fix_pass_comment.py b/tests/unit/workflow/nodes/test_local_review_fix_pass_comment.py index f4c6f3e9..1bc89536 100644 --- a/tests/unit/workflow/nodes/test_local_review_fix_pass_comment.py +++ b/tests/unit/workflow/nodes/test_local_review_fix_pass_comment.py @@ -75,9 +75,7 @@ async def test_posts_fix_pass_comment_on_second_pass(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -111,9 +109,7 @@ async def test_posts_fix_pass_comment_on_third_pass(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -144,9 +140,7 @@ async def test_posts_fix_pass_comment_on_fifth_pass(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -177,9 +171,7 @@ async def test_no_fix_pass_comment_on_first_pass(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -249,9 +241,7 @@ async def test_fix_pass_comment_posted_after_workspace_check(self): with ( patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() result = await local_review_changes(state) @@ -277,9 +267,7 @@ async def test_fix_pass_comment_posted_before_max_attempts_check(self): with ( patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() result = await local_review_changes(state) @@ -317,9 +305,7 @@ async def test_fix_pass_comment_uses_correct_ticket_key(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -350,9 +336,7 @@ async def test_fix_pass_comment_increments_correctly_across_retries(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() result = await local_review_changes(state) diff --git a/tests/unit/workflow/nodes/test_local_review_pass_tracking_errors.py b/tests/unit/workflow/nodes/test_local_review_pass_tracking_errors.py index 2d03b252..c6fc5f3e 100644 --- a/tests/unit/workflow/nodes/test_local_review_pass_tracking_errors.py +++ b/tests/unit/workflow/nodes/test_local_review_pass_tracking_errors.py @@ -1,7 +1,6 @@ """Unit tests for defensive pass number tracking error handling in local_reviewer.py.""" import logging -from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -87,13 +86,14 @@ async def test_none_pass_number_posts_generic_comment(self, caplog): mock_result.stderr = "" mock_runner.run = AsyncMock(return_value=mock_result) - with patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), \ - patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), \ - patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, \ - patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post, \ - caplog.at_level(logging.WARNING): - + with ( + patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), + patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), + patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post, + caplog.at_level(logging.WARNING), + ): mock_git_instance = MagicMock() mock_git_instance.has_uncommitted_changes.return_value = False mock_git_ops.return_value = mock_git_instance @@ -136,12 +136,13 @@ async def test_workflow_continues_when_pass_number_unavailable(self): mock_result.stderr = "" mock_runner.run = AsyncMock(return_value=mock_result) - with patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), \ - patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), \ - patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, \ - patch("forge.workflow.nodes.local_reviewer.post_status_comment"): - + with ( + patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), + patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), + patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, + patch("forge.workflow.nodes.local_reviewer.post_status_comment"), + ): mock_git_instance = MagicMock() mock_git_instance.has_uncommitted_changes.return_value = False mock_git_ops.return_value = mock_git_instance @@ -178,13 +179,14 @@ async def test_negative_pass_number_detected_and_logged(self, caplog): mock_result.stderr = "" mock_runner.run = AsyncMock(return_value=mock_result) - with patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), \ - patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), \ - patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, \ - patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post, \ - caplog.at_level(logging.WARNING): - + with ( + patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), + patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), + patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post, + caplog.at_level(logging.WARNING), + ): mock_git_instance = MagicMock() mock_git_instance.has_uncommitted_changes.return_value = False mock_git_ops.return_value = mock_git_instance @@ -228,13 +230,14 @@ async def test_non_integer_pass_number_detected_and_logged(self, caplog): mock_result.stderr = "" mock_runner.run = AsyncMock(return_value=mock_result) - with patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), \ - patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), \ - patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, \ - patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post, \ - caplog.at_level(logging.WARNING): - + with ( + patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), + patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), + patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post, + caplog.at_level(logging.WARNING), + ): mock_git_instance = MagicMock() mock_git_instance.has_uncommitted_changes.return_value = False mock_git_ops.return_value = mock_git_instance @@ -275,13 +278,14 @@ async def test_zero_pass_number_rejected_with_generic_comment(self, caplog): mock_result.stderr = "" mock_runner.run = AsyncMock(return_value=mock_result) - with patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), \ - patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), \ - patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, \ - patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post, \ - caplog.at_level(logging.WARNING): - + with ( + patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), + patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), + patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post, + caplog.at_level(logging.WARNING), + ): mock_git_instance = MagicMock() mock_git_instance.has_uncommitted_changes.return_value = False mock_git_ops.return_value = mock_git_instance @@ -325,13 +329,14 @@ async def test_pass_one_logs_info_message(self, caplog): mock_result.stderr = "" mock_runner.run = AsyncMock(return_value=mock_result) - with patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), \ - patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), \ - patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, \ - patch("forge.workflow.nodes.local_reviewer.post_status_comment"), \ - caplog.at_level(logging.INFO): - + with ( + patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), + patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), + patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, + patch("forge.workflow.nodes.local_reviewer.post_status_comment"), + caplog.at_level(logging.INFO), + ): mock_git_instance = MagicMock() mock_git_instance.has_uncommitted_changes.return_value = False mock_git_ops.return_value = mock_git_instance @@ -362,13 +367,14 @@ async def test_pass_two_logs_info_message(self, caplog): mock_result.stderr = "" mock_runner.run = AsyncMock(return_value=mock_result) - with patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), \ - patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), \ - patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, \ - patch("forge.workflow.nodes.local_reviewer.post_status_comment"), \ - caplog.at_level(logging.INFO): - + with ( + patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), + patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), + patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, + patch("forge.workflow.nodes.local_reviewer.post_status_comment"), + caplog.at_level(logging.INFO), + ): mock_git_instance = MagicMock() mock_git_instance.has_uncommitted_changes.return_value = False mock_git_ops.return_value = mock_git_instance @@ -399,13 +405,14 @@ async def test_pass_five_logs_info_message(self, caplog): mock_result.stderr = "" mock_runner.run = AsyncMock(return_value=mock_result) - with patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), \ - patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), \ - patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, \ - patch("forge.workflow.nodes.local_reviewer.post_status_comment"), \ - caplog.at_level(logging.INFO): - + with ( + patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), + patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), + patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, + patch("forge.workflow.nodes.local_reviewer.post_status_comment"), + caplog.at_level(logging.INFO), + ): mock_git_instance = MagicMock() mock_git_instance.has_uncommitted_changes.return_value = False mock_git_ops.return_value = mock_git_instance @@ -440,13 +447,14 @@ async def test_warning_log_includes_ticket_key(self, caplog): mock_result.stderr = "" mock_runner.run = AsyncMock(return_value=mock_result) - with patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), \ - patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), \ - patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, \ - patch("forge.workflow.nodes.local_reviewer.post_status_comment"), \ - caplog.at_level(logging.WARNING): - + with ( + patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), + patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), + patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, + patch("forge.workflow.nodes.local_reviewer.post_status_comment"), + caplog.at_level(logging.WARNING), + ): mock_git_instance = MagicMock() mock_git_instance.has_uncommitted_changes.return_value = False mock_git_ops.return_value = mock_git_instance @@ -479,13 +487,14 @@ async def test_warning_log_includes_raw_value_diagnostic(self, caplog): mock_result.stderr = "" mock_runner.run = AsyncMock(return_value=mock_result) - with patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), \ - patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), \ - patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, \ - patch("forge.workflow.nodes.local_reviewer.post_status_comment"), \ - caplog.at_level(logging.WARNING): - + with ( + patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), + patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), + patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, + patch("forge.workflow.nodes.local_reviewer.post_status_comment"), + caplog.at_level(logging.WARNING), + ): mock_git_instance = MagicMock() mock_git_instance.has_uncommitted_changes.return_value = False mock_git_ops.return_value = mock_git_instance @@ -520,12 +529,13 @@ async def test_pass_number_increments_correctly_after_retry(self): mock_result.stderr = "" mock_runner.run = AsyncMock(return_value=mock_result) - with patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), \ - patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), \ - patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, \ - patch("forge.workflow.nodes.local_reviewer.post_status_comment"): - + with ( + patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), + patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), + patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, + patch("forge.workflow.nodes.local_reviewer.post_status_comment"), + ): mock_git_instance = MagicMock() mock_git_instance.has_uncommitted_changes.return_value = False mock_git_ops.return_value = mock_git_instance @@ -557,12 +567,13 @@ async def test_pass_number_recovers_from_none_and_increments(self): mock_result.stderr = "" mock_runner.run = AsyncMock(return_value=mock_result) - with patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), \ - patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), \ - patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), \ - patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, \ - patch("forge.workflow.nodes.local_reviewer.post_status_comment"): - + with ( + patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), + patch("forge.workflow.nodes.local_reviewer.load_prompt", return_value="test prompt"), + patch("forge.workflow.nodes.local_reviewer.GitOperations") as mock_git_ops, + patch("forge.workflow.nodes.local_reviewer.post_status_comment"), + ): mock_git_instance = MagicMock() mock_git_instance.has_uncommitted_changes.return_value = False mock_git_ops.return_value = mock_git_instance diff --git a/tests/unit/workflow/nodes/test_local_review_status_comments_comprehensive.py b/tests/unit/workflow/nodes/test_local_review_status_comments_comprehensive.py index fc12d529..f16b9dc6 100644 --- a/tests/unit/workflow/nodes/test_local_review_status_comments_comprehensive.py +++ b/tests/unit/workflow/nodes/test_local_review_status_comments_comprehensive.py @@ -59,7 +59,7 @@ def create_mock_git_operations(has_changes=False): class TestPassNumberOneCommentPosting: """Tests verifying initial comment posts only when pass_number == 1. - + Acceptance Criteria: Unit tests verify initial comment posts only when pass_number == 1 """ @@ -82,9 +82,7 @@ async def test_posts_initial_comment_when_pass_number_equals_one(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -116,16 +114,16 @@ async def test_no_initial_comment_when_pass_number_equals_two(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) # Verify initial comment (with 🔍) was NOT posted for call in mock_post_status.call_args_list: - assert "🔍" not in str(call), "Initial comment should not be posted when pass_number > 1" + assert "🔍" not in str(call), ( + "Initial comment should not be posted when pass_number > 1" + ) @pytest.mark.asyncio async def test_no_initial_comment_when_pass_number_greater_than_one(self): @@ -146,9 +144,7 @@ async def test_no_initial_comment_when_pass_number_greater_than_one(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -162,7 +158,7 @@ async def test_no_initial_comment_when_pass_number_greater_than_one(self): class TestPassNumberGreaterThanOneCommentPosting: """Tests verifying fix comments post only when pass_number > 1. - + Acceptance Criteria: Unit tests verify fix comments post only when pass_number > 1 """ @@ -185,9 +181,7 @@ async def test_posts_fix_comment_when_pass_number_equals_two(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -219,9 +213,7 @@ async def test_posts_fix_comment_when_pass_number_greater_than_two(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -250,9 +242,7 @@ async def test_no_fix_comment_when_pass_number_equals_one(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -264,8 +254,8 @@ async def test_no_fix_comment_when_pass_number_equals_one(self): class TestCorrectPassNumberInCommentText: """Tests verifying correct pass number appears in comment text. - - Acceptance Criteria: Unit tests verify correct pass number appears in comment text + + Acceptance Criteria: Unit tests verify correct pass number appears in comment text for passes 2, 3, 4, 5+ """ @@ -288,9 +278,7 @@ async def test_comment_shows_pass_two_correctly(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -321,9 +309,7 @@ async def test_comment_shows_pass_three_correctly(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -354,9 +340,7 @@ async def test_comment_shows_pass_four_correctly(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -388,9 +372,7 @@ async def test_comment_shows_pass_five_plus_correctly(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -421,9 +403,7 @@ async def test_comment_shows_high_pass_number_correctly(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -438,7 +418,7 @@ async def test_comment_shows_high_pass_number_correctly(self): class TestGracefulHandlingWhenPassNumberUnavailable: """Tests verifying graceful handling when pass_number unavailable. - + Acceptance Criteria: Unit tests verify graceful handling when pass_number unavailable """ @@ -463,9 +443,7 @@ async def test_defaults_to_pass_one_when_pass_number_missing(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -498,9 +476,7 @@ async def test_workflow_completes_successfully_without_pass_number(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() result = await local_review_changes(state) @@ -529,12 +505,10 @@ async def test_no_error_when_pass_number_none(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() - + # Should not raise exception try: result = await local_review_changes(state) @@ -562,12 +536,10 @@ async def test_handles_pass_number_zero_gracefully(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() - + # Should not raise exception result = await local_review_changes(state) @@ -579,7 +551,9 @@ async def test_handles_pass_number_zero_gracefully(self): comment_args = mock_post_status.call_args[0] assert comment_args[0] == mock_jira # First arg is jira client assert comment_args[1] == "FEAT-503" # Second arg is ticket key - assert "🔧 Local review found issues, applying fixes." in comment_args[2] # Third arg is message + assert ( + "🔧 Local review found issues, applying fixes." in comment_args[2] + ) # Third arg is message class TestIntegrationWithReviewFlow: @@ -647,9 +621,7 @@ async def test_comment_posted_to_correct_ticket(self): patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.local_reviewer.ContainerRunner", return_value=mock_runner), patch("forge.workflow.nodes.local_reviewer.GitOperations", return_value=mock_git), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): mock_post_status.return_value = AsyncMock() await local_review_changes(state) @@ -669,9 +641,7 @@ async def test_no_comment_when_workspace_missing(self): with ( patch("forge.workflow.nodes.local_reviewer.JiraClient", return_value=mock_jira), - patch( - "forge.workflow.nodes.local_reviewer.post_status_comment" - ) as mock_post_status, + patch("forge.workflow.nodes.local_reviewer.post_status_comment") as mock_post_status, ): result = await local_review_changes(state) diff --git a/tests/unit/workflow/nodes/test_qa_handler.py b/tests/unit/workflow/nodes/test_qa_handler.py index ea4a1bf9..a7da90a6 100644 --- a/tests/unit/workflow/nodes/test_qa_handler.py +++ b/tests/unit/workflow/nodes/test_qa_handler.py @@ -20,7 +20,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.""" @@ -599,7 +601,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_rca_analysis.py b/tests/unit/workflow/nodes/test_rca_analysis.py index f25f9326..b514d4cc 100644 --- a/tests/unit/workflow/nodes/test_rca_analysis.py +++ b/tests/unit/workflow/nodes/test_rca_analysis.py @@ -200,7 +200,9 @@ async def run(self, workspace_path, task_description="", **_kwargs): with ( patch("forge.workflow.nodes.rca_analysis.JiraClient", return_value=mock_jira), - patch("forge.workflow.nodes.rca_analysis.ContainerRunner", return_value=_CapturingRunner()), + patch( + "forge.workflow.nodes.rca_analysis.ContainerRunner", return_value=_CapturingRunner() + ), ): result = await analyze_bug(base_bug_state) @@ -407,7 +409,9 @@ async def run(self, workspace_path, task_description="", **_kwargs): with ( patch("forge.workflow.nodes.rca_analysis.JiraClient", return_value=mock_jira), - patch("forge.workflow.nodes.rca_analysis.ContainerRunner", return_value=_CapturingRunner()), + patch( + "forge.workflow.nodes.rca_analysis.ContainerRunner", return_value=_CapturingRunner() + ), ): await reflect_rca(rca_state) @@ -438,12 +442,16 @@ class _HistoryRunner: async def run(self, workspace_path, **_kwargs): history_dir = workspace_path / ".forge" / "history" history_dir.mkdir(parents=True, exist_ok=True) - (history_dir / "BUG-123-reflect.json").write_text(json.dumps({ - "messages": [ - {"role": "human", "content": "Review this RCA"}, - {"role": "ai", "content": "VALID"}, - ], - })) + (history_dir / "BUG-123-reflect.json").write_text( + json.dumps( + { + "messages": [ + {"role": "human", "content": "Review this RCA"}, + {"role": "ai", "content": "VALID"}, + ], + } + ) + ) result = MagicMock() result.success = True result.exit_code = 0 @@ -457,7 +465,9 @@ async def run(self, workspace_path, **_kwargs): with ( patch("forge.workflow.nodes.rca_analysis.JiraClient", return_value=mock_jira), - patch("forge.workflow.nodes.rca_analysis.ContainerRunner", return_value=_HistoryRunner()), + patch( + "forge.workflow.nodes.rca_analysis.ContainerRunner", return_value=_HistoryRunner() + ), ): result = await reflect_rca(rca_state) diff --git a/tests/unit/workflow/nodes/test_rca_option_gate.py b/tests/unit/workflow/nodes/test_rca_option_gate.py index bbad9a39..7efe0ee0 100644 --- a/tests/unit/workflow/nodes/test_rca_option_gate.py +++ b/tests/unit/workflow/nodes/test_rca_option_gate.py @@ -140,7 +140,7 @@ async def test_truncation_preserves_paragraph_boundary(self): """Truncation happens at the last \\n\\n before the limit, not mid-sentence.""" # Build rca_content with paragraphs separated by \n\n paragraph = "Word " * 100 # ~500 chars per paragraph - rca = ("\n\n".join([paragraph] * 60)) # ~30k chars + rca = "\n\n".join([paragraph] * 60) # ~30k chars state = make_rca_option_state(rca_content=rca) mock_jira = _make_mock_jira() diff --git a/tests/unit/workflow/nodes/test_task_generation.py b/tests/unit/workflow/nodes/test_task_generation.py index 95c5a813..8e5ef1bc 100644 --- a/tests/unit/workflow/nodes/test_task_generation.py +++ b/tests/unit/workflow/nodes/test_task_generation.py @@ -5,6 +5,7 @@ import pytest from forge.integrations.jira.models import JiraIssue +from forge.models.workflow import ForgeLabel from forge.workflow.nodes.task_generation import ( _generate_tasks_for_epic, _parse_tasks_response, @@ -12,6 +13,7 @@ regenerate_all_tasks, regenerate_epic_tasks, ) +from forge.workflow.utils.draft_manager import DraftManager @pytest.fixture @@ -24,6 +26,7 @@ def base_state(): "task_keys": [], "tasks_by_repo": {}, "retry_count": 0, + "yolo_mode": True, } @@ -585,3 +588,170 @@ async def test_orphaned_task_with_none_parent_logged_as_warning(self, base_state r for r in caplog.records if "TASK-100" in r.message and "parent" in r.message.lower() ] assert orphan_warnings, "Expected a warning about the orphaned task TASK-100" + + +class TestTaskGenerationDraftReview: + """Tests for the non-YOLO draft review gate flow in generate_tasks.""" + + @pytest.mark.asyncio + async def test_generate_tasks_draft_review_flow_success( + self, base_state, mock_parent_issue, mock_epic_issue, mock_tasks_data + ): + """When yolo_mode is False, generates tasks into a draft JSON, deletes old attachments, saves new one, posts comment, and pauses.""" + state = {**base_state, "yolo_mode": False} + + with ( + patch("forge.workflow.nodes.task_generation.JiraClient") as MockJira, + patch("forge.workflow.nodes.task_generation.ForgeAgent") as MockAgent, + patch("forge.workflow.nodes.task_generation.DraftManager") as MockDraftManager, + patch("forge.workflow.nodes.task_generation.post_status_comment"), + patch( + "forge.workflow.nodes.task_generation._generate_tasks_for_epic", + new_callable=AsyncMock, + return_value=mock_tasks_data, + ), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(side_effect=[mock_parent_issue, mock_epic_issue]) + mock_jira.get_labels = AsyncMock(return_value=[]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + + MockDraftManager.delete_draft_attachment = AsyncMock() + MockDraftManager.save_draft_attachment = AsyncMock() + MockDraftManager.format_review_comment.side_effect = DraftManager.format_review_comment + + result = await generate_tasks(state) + + # 1. Verify DraftManager deleted any existing forge-tasks-draft.json first + MockDraftManager.delete_draft_attachment.assert_called_once_with( + mock_jira, "MYPROJ-1", "forge-tasks-draft.json" + ) + + # 2. Verify DraftManager saved the new draft + MockDraftManager.save_draft_attachment.assert_called_once() + saved_draft = MockDraftManager.save_draft_attachment.call_args[0][2] + assert saved_draft.phase == "tasks" + assert len(saved_draft.items) == 1 + assert saved_draft.items[0].summary == "Task One" + assert saved_draft.items[0].description == "Do the first thing." + assert saved_draft.items[0].repo == "acme/backend" + + # 3. Verify formatted comment posted + assert mock_jira.add_comment.call_count == 1 + comment_text = mock_jira.add_comment.call_args[0][1] + assert "### 📋 Proposed Tasks Draft" in comment_text + assert "Task One" in comment_text + assert "acme/backend" in comment_text + + # 4. Verify workflow label updated to TASK_PENDING + mock_jira.set_workflow_label.assert_called_once_with("MYPROJ-1", ForgeLabel.TASK_PENDING) + + # 5. Verify state transitions to task_approval_gate and pauses + assert result["current_node"] == "task_approval_gate" + assert result["is_paused"] is True + assert result["task_keys"] == [] + + @pytest.mark.asyncio + async def test_generate_tasks_draft_review_truncation_limits( + self, base_state, mock_parent_issue, mock_epic_issue + ): + """When the item list has > 15 elements, comment falls back to a condensed table.""" + state = {**base_state, "yolo_mode": False} + + # Mock 16 items + many_tasks_data = [ + {"summary": f"Task {i}", "description": f"Desc {i}", "repo": f"acme/repo-{i}"} + for i in range(1, 17) + ] + + with ( + patch("forge.workflow.nodes.task_generation.JiraClient") as MockJira, + patch("forge.workflow.nodes.task_generation.ForgeAgent") as MockAgent, + patch("forge.workflow.nodes.task_generation.DraftManager") as MockDraftManager, + patch("forge.workflow.nodes.task_generation.post_status_comment"), + patch( + "forge.workflow.nodes.task_generation._generate_tasks_for_epic", + new_callable=AsyncMock, + return_value=many_tasks_data, + ), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(side_effect=[mock_parent_issue, mock_epic_issue]) + mock_jira.get_labels = AsyncMock(return_value=[]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + + MockDraftManager.delete_draft_attachment = AsyncMock() + MockDraftManager.save_draft_attachment = AsyncMock() + MockDraftManager.format_review_comment.side_effect = DraftManager.format_review_comment + + await generate_tasks(state) + + # Verify comment is in condensed table format + assert mock_jira.add_comment.call_count == 1 + comment_text = mock_jira.add_comment.call_args[0][1] + assert "### 📋 Proposed Tasks Draft (Condensed)" in comment_text + assert "Warning" in comment_text + assert "forge-tasks-draft.json" in comment_text + # Condensed table should only show IDs, summaries, and target repos + # Detailed descriptions/plans (like Desc 1) should NOT be in the comment + assert "Desc 1" not in comment_text + assert "Task 1" in comment_text + assert "acme/repo-1" in comment_text + + @pytest.mark.asyncio + async def test_generate_tasks_draft_review_truncation_characters( + self, base_state, mock_parent_issue, mock_epic_issue + ): + """When comment exceeds 32,767 characters, comment falls back to a condensed table.""" + state = {**base_state, "yolo_mode": False} + + # Mock huge description to exceed character limit + huge_tasks_data = [ + {"summary": "Task One", "description": "A" * 35000, "repo": "acme/backend"} + ] + + with ( + patch("forge.workflow.nodes.task_generation.JiraClient") as MockJira, + patch("forge.workflow.nodes.task_generation.ForgeAgent") as MockAgent, + patch("forge.workflow.nodes.task_generation.DraftManager") as MockDraftManager, + patch("forge.workflow.nodes.task_generation.post_status_comment"), + patch( + "forge.workflow.nodes.task_generation._generate_tasks_for_epic", + new_callable=AsyncMock, + return_value=huge_tasks_data, + ), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(side_effect=[mock_parent_issue, mock_epic_issue]) + mock_jira.get_labels = AsyncMock(return_value=[]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + + MockDraftManager.delete_draft_attachment = AsyncMock() + MockDraftManager.save_draft_attachment = AsyncMock() + MockDraftManager.format_review_comment.side_effect = DraftManager.format_review_comment + + await generate_tasks(state) + + # Verify comment is in condensed table format due to length + assert mock_jira.add_comment.call_count == 1 + comment_text = mock_jira.add_comment.call_args[0][1] + assert "Warning" in comment_text + assert "forge-tasks-draft.json" in comment_text + assert "A" * 35000 not in comment_text + assert "Task One" in comment_text + assert "acme/backend" in comment_text diff --git a/tests/unit/workflow/nodes/test_task_takeover_planning.py b/tests/unit/workflow/nodes/test_task_takeover_planning.py index e2055846..64422346 100644 --- a/tests/unit/workflow/nodes/test_task_takeover_planning.py +++ b/tests/unit/workflow/nodes/test_task_takeover_planning.py @@ -110,7 +110,9 @@ async def test_generate_plan_success(self, base_task_state: TaskTakeoverState) - ): result = await generate_plan(base_task_state) - assert result["plan_content"] == "## Plan\n\nTask Takeover Plan details.\n\nrepo:owner/project" + assert ( + result["plan_content"] == "## Plan\n\nTask Takeover Plan details.\n\nrepo:owner/project" + ) assert result["current_repo"] == "owner/project" assert result["repos_to_process"] == ["owner/project"] assert result["current_node"] == "task_plan_approval_gate" @@ -208,7 +210,10 @@ async def test_regenerate_plan_with_feedback(self, base_task_state: TaskTakeover ): result = await generate_plan(state) - assert result["plan_content"] == "## Plan\n\nNew Plan content with logging.\n\nrepo:owner/project" + assert ( + result["plan_content"] + == "## Plan\n\nNew Plan content with logging.\n\nrepo:owner/project" + ) assert result["revision_requested"] is False assert result["feedback_comment"] is None assert result["current_node"] == "task_plan_approval_gate" diff --git a/tests/unit/workflow/nodes/test_task_takeover_triage.py b/tests/unit/workflow/nodes/test_task_takeover_triage.py index 242a7348..2e43365a 100644 --- a/tests/unit/workflow/nodes/test_task_takeover_triage.py +++ b/tests/unit/workflow/nodes/test_task_takeover_triage.py @@ -69,9 +69,7 @@ def mock_agent_sufficient() -> MagicMock: def mock_agent_missing_fields() -> MagicMock: """ForgeAgent that returns a JSON list of missing fields.""" agent = MagicMock() - agent.run_task = AsyncMock( - return_value='["Problem Statement", "Acceptance Criteria"]' - ) + agent.run_task = AsyncMock(return_value='["Problem Statement", "Acceptance Criteria"]') agent.close = AsyncMock() return agent @@ -90,9 +88,7 @@ async def test_sets_triage_passed_true( from forge.workflow.nodes.task_takeover_triage import triage_task with ( - patch( - "forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.task_takeover_triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -130,9 +126,7 @@ async def mock_run_task(*_args: Any, **_kwargs: Any) -> str: mock_agent_sufficient.run_task.side_effect = mock_run_task with ( - patch( - "forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.task_takeover_triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -154,9 +148,7 @@ async def test_acknowledgement_comment_suppressed_on_resume( from forge.workflow.nodes.task_takeover_triage import triage_task with ( - patch( - "forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.task_takeover_triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -188,9 +180,7 @@ async def test_resume_with_complete_ticket_consumes_revision_signal( } with ( - patch( - "forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.task_takeover_triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -223,9 +213,7 @@ async def test_sufficient_ticket_sets_inferred_repo( mock_jira.get_project_default_repo = AsyncMock(return_value="openshift/installer") with ( - patch( - "forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.task_takeover_triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -252,9 +240,7 @@ async def test_sets_triage_passed_false( from forge.workflow.nodes.task_takeover_triage import triage_task with ( - patch( - "forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.task_takeover_triage.ForgeAgent", return_value=mock_agent_missing_fields, @@ -280,9 +266,7 @@ async def test_applies_triage_pending_label_and_posts_comment( from forge.workflow.nodes.task_takeover_triage import triage_task with ( - patch( - "forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.task_takeover_triage.ForgeAgent", return_value=mock_agent_missing_fields, @@ -310,9 +294,7 @@ async def test_escalates_to_blocked_on_max_retries(self, mock_jira: MagicMock) - state = make_task_state(retry_count=3) with ( - patch( - "forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.task_takeover_triage.JiraClient", return_value=mock_jira), ): result = await triage_task(state) diff --git a/tests/unit/workflow/nodes/test_trace_context_enrichment.py b/tests/unit/workflow/nodes/test_trace_context_enrichment.py index be31f9aa..9ced1870 100644 --- a/tests/unit/workflow/nodes/test_trace_context_enrichment.py +++ b/tests/unit/workflow/nodes/test_trace_context_enrichment.py @@ -354,9 +354,7 @@ async def test_update_single_epic_passes_trace_fields(self) -> None: mock_jira = MagicMock() mock_jira.close = AsyncMock() - mock_jira.get_issue = AsyncMock( - return_value=MagicMock(description="Original epic") - ) + mock_jira.get_issue = AsyncMock(return_value=MagicMock(description="Original epic")) mock_jira.update_description = AsyncMock() mock_jira.add_comment = AsyncMock() @@ -404,9 +402,7 @@ async def test_update_single_task_passes_trace_fields(self) -> None: mock_jira = MagicMock() mock_jira.close = AsyncMock() - mock_jira.get_issue = AsyncMock( - return_value=MagicMock(description="Original task") - ) + mock_jira.get_issue = AsyncMock(return_value=MagicMock(description="Original task")) mock_jira.update_description = AsyncMock() mock_jira.add_comment = AsyncMock() diff --git a/tests/unit/workflow/nodes/test_triage.py b/tests/unit/workflow/nodes/test_triage.py index c38602f9..ef564208 100644 --- a/tests/unit/workflow/nodes/test_triage.py +++ b/tests/unit/workflow/nodes/test_triage.py @@ -77,9 +77,7 @@ def mock_agent_sufficient(): def mock_agent_missing_fields(): """ForgeAgent that returns a JSON list of missing fields.""" agent = MagicMock() - agent.run_task = AsyncMock( - return_value='["steps_to_reproduce", "environment"]' - ) + agent.run_task = AsyncMock(return_value='["steps_to_reproduce", "environment"]') agent.close = AsyncMock() return agent @@ -95,9 +93,7 @@ async def test_sets_triage_passed_true( from forge.workflow.nodes.triage import triage_check with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -114,9 +110,7 @@ async def test_missing_fields_empty( from forge.workflow.nodes.triage import triage_check with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -133,9 +127,7 @@ async def test_no_triage_pending_label_set( from forge.workflow.nodes.triage import triage_check with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -160,9 +152,7 @@ async def test_acknowledgement_comment_posted_first( side_effect=lambda *_a, **_k: call_order.append("agent") or "sufficient" ) with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -185,9 +175,7 @@ async def test_acknowledgement_comment_suppressed_on_resume( triage_missing_fields=["steps_to_reproduce"], ) with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -207,9 +195,7 @@ async def test_acknowledgement_comment_content( from forge.workflow.nodes.triage import triage_check with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -235,9 +221,7 @@ async def test_sets_triage_passed_false( from forge.workflow.nodes.triage import triage_check with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_missing_fields, @@ -254,9 +238,7 @@ async def test_missing_fields_populated( from forge.workflow.nodes.triage import triage_check with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_missing_fields, @@ -274,9 +256,7 @@ async def test_targeted_comment_posted( from forge.workflow.nodes.triage import triage_check with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_missing_fields, @@ -287,10 +267,7 @@ async def test_targeted_comment_posted( assert mock_jira.add_comment.call_count >= 2 last_comment = mock_jira.add_comment.call_args_list[-1].args[1] assert "starting with `!`" in last_comment - assert ( - "steps_to_reproduce" in last_comment - or "steps to reproduce" in last_comment.lower() - ) + assert "steps_to_reproduce" in last_comment or "steps to reproduce" in last_comment.lower() @pytest.mark.asyncio async def test_triage_pending_label_set( @@ -300,9 +277,7 @@ async def test_triage_pending_label_set( from forge.workflow.nodes.triage import triage_check with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_missing_fields, @@ -321,9 +296,7 @@ async def test_current_node_set_to_triage_gate( from forge.workflow.nodes.triage import triage_check with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_missing_fields, @@ -337,9 +310,7 @@ class TestTriageCheckResume: """triage_check re-evaluates on resume after reporter updates ticket.""" @pytest.mark.asyncio - async def test_resume_with_complete_ticket_passes( - self, mock_jira, mock_agent_sufficient - ): + async def test_resume_with_complete_ticket_passes(self, mock_jira, mock_agent_sufficient): """On resume, if ticket now has all fields, triage_passed=True.""" from forge.workflow.nodes.triage import triage_check @@ -350,9 +321,7 @@ async def test_resume_with_complete_ticket_passes( triage_missing_fields=["steps_to_reproduce"], ) with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -378,9 +347,7 @@ async def test_resume_with_complete_ticket_consumes_revision_signal( is_question=True, ) with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_sufficient, @@ -396,9 +363,7 @@ async def test_resume_with_complete_ticket_consumes_revision_signal( assert result["feedback_comment"] is None @pytest.mark.asyncio - async def test_resume_still_missing_reposts_comment( - self, mock_jira, mock_agent_missing_fields - ): + async def test_resume_still_missing_reposts_comment(self, mock_jira, mock_agent_missing_fields): """On resume, still-missing fields cause a fresh targeted comment.""" from forge.workflow.nodes.triage import triage_check @@ -409,9 +374,7 @@ async def test_resume_still_missing_reposts_comment( triage_missing_fields=["steps_to_reproduce"], ) with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), patch( "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent_missing_fields, @@ -427,9 +390,7 @@ class TestTriageCheckErrorHandling: """triage_check retries on failure and escalates after 3 failures.""" @pytest.mark.asyncio - async def test_failure_increments_retry_count( - self, incomplete_ticket_state, mock_jira - ): + async def test_failure_increments_retry_count(self, incomplete_ticket_state, mock_jira): """Node failure increments retry_count.""" from forge.workflow.nodes.triage import triage_check @@ -438,20 +399,14 @@ async def test_failure_increments_retry_count( mock_agent.close = AsyncMock() incomplete_ticket_state["retry_count"] = 1 with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), - patch( - "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent), ): result = await triage_check(incomplete_ticket_state) assert result["retry_count"] == 2 @pytest.mark.asyncio - async def test_after_3_failures_escalates_blocked( - self, incomplete_ticket_state, mock_jira - ): + async def test_after_3_failures_escalates_blocked(self, incomplete_ticket_state, mock_jira): """After 3 consecutive failures (retry_count already at max), routes to escalate_blocked.""" from forge.workflow.nodes.triage import triage_check @@ -460,12 +415,8 @@ async def test_after_3_failures_escalates_blocked( mock_agent.close = AsyncMock() incomplete_ticket_state["retry_count"] = 3 with ( - patch( - "forge.workflow.nodes.triage.JiraClient", return_value=mock_jira - ), - patch( - "forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent - ), + patch("forge.workflow.nodes.triage.JiraClient", return_value=mock_jira), + patch("forge.workflow.nodes.triage.ForgeAgent", return_value=mock_agent), ): result = await triage_check(incomplete_ticket_state) assert result["current_node"] == "escalate_blocked" diff --git a/tests/unit/workflow/nodes/test_workspace_setup.py b/tests/unit/workflow/nodes/test_workspace_setup.py index ca06cd7d..0a27ed61 100644 --- a/tests/unit/workflow/nodes/test_workspace_setup.py +++ b/tests/unit/workflow/nodes/test_workspace_setup.py @@ -293,13 +293,9 @@ 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_github): """Implementation must not start without its durable backup remote.""" - mock_workspace_github.get_or_create_fork.side_effect = RuntimeError( - "fork creation denied" - ) + mock_workspace_github.get_or_create_fork.side_effect = RuntimeError("fork creation denied") mock_jira = create_mock_jira_client() mock_manager, _ = create_mock_workspace_manager() mock_git = create_mock_git_operations() @@ -329,9 +325,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_github): mock_jira = create_mock_jira_client() mock_manager, _ = create_mock_workspace_manager() mock_git = create_mock_git_operations() @@ -352,9 +346,7 @@ async def test_creates_fork_remote_before_implementation( ): result = await setup_workspace(state) - mock_workspace_github.get_or_create_fork.assert_awaited_once_with( - "upstream", "repo" - ) + 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" ) @@ -389,18 +381,14 @@ 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_github.get_or_create_fork.assert_awaited_once_with("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_github): mock_jira = create_mock_jira_client() mock_manager, mock_workspace = create_mock_workspace_manager() mock_git = create_mock_git_operations() @@ -428,9 +416,7 @@ async def test_existing_fork_branch_is_checked_out_without_push( mock_git.remote_branch_exists.assert_called_once_with( mock_workspace.branch_name, remote="fork" ) - mock_git.checkout_branch.assert_called_once_with( - mock_workspace.branch_name, remote="fork" - ) + mock_git.checkout_branch.assert_called_once_with(mock_workspace.branch_name, remote="fork") mock_git.create_branch.assert_not_called() mock_git.push_to_fork.assert_not_called() assert result["current_node"] == "implementation" diff --git a/tests/unit/workflow/test_base.py b/tests/unit/workflow/test_base.py index 4df75da1..dfd66aba 100644 --- a/tests/unit/workflow/test_base.py +++ b/tests/unit/workflow/test_base.py @@ -131,6 +131,7 @@ class ConcreteWorkflow(BaseWorkflow): @property def state_schema(self): from forge.workflow.base import BaseState + return BaseState def matches(self, ticket_type, labels, event): diff --git a/tests/unit/workflow/test_ci_gate_skip.py b/tests/unit/workflow/test_ci_gate_skip.py index 89da27a2..fbd3c1bb 100644 --- a/tests/unit/workflow/test_ci_gate_skip.py +++ b/tests/unit/workflow/test_ci_gate_skip.py @@ -3,11 +3,11 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from tests.fixtures.workflow_states import make_workflow_state from forge.models.events import EventSource from forge.orchestrator.worker import OrchestratorWorker from forge.queue.models import QueueMessage +from tests.fixtures.workflow_states import make_workflow_state # ── Helpers ─────────────────────────────────────────────────────────────────── @@ -85,16 +85,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 +105,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 +116,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 +127,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 +153,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 +164,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 +178,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,9 +189,7 @@ 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( @@ -225,7 +214,6 @@ 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.""" @@ -239,8 +227,10 @@ async def test_posts_github_reply_and_jira_comment(self): 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.GitHubClient", return_value=mock_github), + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + ): await worker._post_skip_gate_feedback( ticket_key="TEST-123", owner="org", @@ -267,8 +257,10 @@ async def test_unskip_posts_different_message(self): 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.GitHubClient", return_value=mock_github), + patch("forge.orchestrator.worker.JiraClient", return_value=mock_jira), + ): await worker._post_skip_gate_feedback( ticket_key="TEST-123", owner="org", @@ -287,7 +279,6 @@ async def test_unskip_posts_different_message(self): 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.""" @@ -301,12 +292,20 @@ async def test_skipped_check_does_not_count_as_failure(self): 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.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): @@ -328,12 +327,20 @@ async def test_all_skipped_checks_plus_pass_routes_to_human_review(self): 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.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): @@ -355,12 +362,16 @@ async def test_skipped_check_not_in_failed_checks(self): 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.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): @@ -383,10 +394,15 @@ async def test_substring_match_is_case_insensitive(self): 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.get_check_runs = AsyncMock( + return_value=[ + { + "name": "Run acceptance tests against OpenStack epoxy", + "status": "completed", + "conclusion": "failure", + }, + ] + ) mock_github.close = AsyncMock() with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): @@ -411,15 +427,20 @@ async def test_tide_is_ignored_as_permanent_pending_check(self): 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.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): @@ -442,12 +463,17 @@ async def test_real_pending_check_still_blocks_evaluation(self): 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.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): @@ -469,9 +495,11 @@ async def test_empty_skipped_checks_behaves_normally(self): 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.get_check_runs = AsyncMock( + return_value=[ + {"name": "unit-tests", "status": "completed", "conclusion": "failure"}, + ] + ) mock_github.close = AsyncMock() with patch("forge.workflow.nodes.ci_evaluator.GitHubClient", return_value=mock_github): diff --git a/tests/unit/workflow/test_cleanup.py b/tests/unit/workflow/test_cleanup.py index a63cceff..25d726b0 100644 --- a/tests/unit/workflow/test_cleanup.py +++ b/tests/unit/workflow/test_cleanup.py @@ -63,6 +63,7 @@ class TestRouteEntryCompleteness: def _route(self, node: str): from forge.workflow.bug.graph import route_entry + return route_entry({"current_node": node}) def test_all_new_pipeline_nodes_mapped(self): @@ -82,9 +83,7 @@ def test_all_new_pipeline_nodes_mapped(self): } for node, expected in new_nodes.items(): result = self._route(node) - assert result == expected, ( - f"route_entry('{node}') = '{result}', expected '{expected}'" - ) + assert result == expected, f"route_entry('{node}') = '{result}', expected '{expected}'" def test_backward_compat_rca_approval_gate(self): """Old rca_approval_gate checkpoint maps to rca_option_gate.""" @@ -93,6 +92,7 @@ def test_backward_compat_rca_approval_gate(self): def test_existing_nodes_still_mapped(self): """All pre-redesign node mappings are preserved.""" from langgraph.graph import END + preserved = { "setup_workspace": "setup_workspace", "implement_bug_fix": "implement_bug_fix", @@ -111,6 +111,4 @@ def test_existing_nodes_still_mapped(self): } for node, expected in preserved.items(): result = self._route(node) - assert result == expected, ( - f"route_entry('{node}') = '{result}', expected '{expected}'" - ) + assert result == expected, f"route_entry('{node}') = '{result}', expected '{expected}'" diff --git a/tests/unit/workflow/test_comment_classifier.py b/tests/unit/workflow/test_comment_classifier.py index 2bfcc7b7..05f120dc 100644 --- a/tests/unit/workflow/test_comment_classifier.py +++ b/tests/unit/workflow/test_comment_classifier.py @@ -92,3 +92,39 @@ def test_whitespace_only_comment_is_informational(self) -> None: """Whitespace-only comments should be informational.""" assert classify_comment(" ") == CommentType.INFORMATIONAL assert classify_comment("\n\t") == CommentType.INFORMATIONAL + + # Command detection tests + def test_command_remove(self) -> None: + """/forge remove command should be classified as a command.""" + assert classify_comment("/forge remove 2") == CommentType.COMMAND + assert classify_comment("/Forge remove abc") == CommentType.COMMAND + + def test_command_exclude(self) -> None: + """/forge exclude command should be classified as a command.""" + assert classify_comment("/forge exclude 3") == CommentType.COMMAND + + def test_command_approve(self) -> None: + """/forge approve command should be classified as a command.""" + assert classify_comment("/forge approve") == CommentType.COMMAND + + def test_command_add(self) -> None: + """/forge add command should be classified as a command.""" + assert classify_comment('/forge add summary="Implement API"') == CommentType.COMMAND + + def test_command_update(self) -> None: + """/forge update command should be classified as a command.""" + assert classify_comment('/forge update 1 summary="New Summary"') == CommentType.COMMAND + + def test_command_case_insensitive_prefix(self) -> None: + """/forge commands should be case-insensitive.""" + assert classify_comment("/FORGE remove 2") == CommentType.COMMAND + assert classify_comment(" /Forge exclude 3") == CommentType.COMMAND + + def test_command_skip_gate_is_ignored_by_classifier(self) -> None: + """skip-gate/unskip-gate are not classified as COMMAND by classify_comment.""" + assert classify_comment("/forge skip-gate build") == CommentType.INFORMATIONAL + assert classify_comment("/forge unskip-gate test") == CommentType.INFORMATIONAL + + def test_command_rebase_is_ignored_by_classifier(self) -> None: + """rebase is not classified as COMMAND by classify_comment.""" + assert classify_comment("/forge rebase") == CommentType.INFORMATIONAL diff --git a/tests/unit/workflow/test_pr_status_comments.py b/tests/unit/workflow/test_pr_status_comments.py index 7a5deaf5..62168a64 100644 --- a/tests/unit/workflow/test_pr_status_comments.py +++ b/tests/unit/workflow/test_pr_status_comments.py @@ -71,7 +71,10 @@ async def test_pr_number_extraction_with_missing_pr_number(self): # Verify fallback message used assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args - assert comment_call[0][1] == "🚀 Pull request created and submitted. Waiting for CI checks to complete." + assert ( + comment_call[0][1] + == "🚀 Pull request created and submitted. Waiting for CI checks to complete." + ) assert "#" not in comment_call[0][1] @pytest.mark.asyncio @@ -93,7 +96,10 @@ async def test_pr_number_extraction_with_malformed_response(self): # Verify fallback message used when key is missing assert mock_jira.add_comment.call_count == 1 comment_call = mock_jira.add_comment.call_args - assert comment_call[0][1] == "🚀 Pull request created and submitted. Waiting for CI checks to complete." + assert ( + comment_call[0][1] + == "🚀 Pull request created and submitted. Waiting for CI checks to complete." + ) class TestPRStatusCommentPosting: @@ -118,7 +124,7 @@ async def test_status_comment_posted_with_pr_number_present(self): # Verify comment posted to correct ticket with correct message mock_jira.add_comment.assert_called_once_with( "TEST-200", - "🚀 Pull request #999 created and submitted. Waiting for CI checks to complete." + "🚀 Pull request #999 created and submitted. Waiting for CI checks to complete.", ) @pytest.mark.asyncio @@ -139,8 +145,7 @@ async def test_status_comment_posted_with_pr_number_absent(self): # Verify fallback comment posted to correct ticket mock_jira.add_comment.assert_called_once_with( - "TEST-201", - "🚀 Pull request created and submitted. Waiting for CI checks to complete." + "TEST-201", "🚀 Pull request created and submitted. Waiting for CI checks to complete." ) @pytest.mark.asyncio @@ -183,10 +188,7 @@ async def test_label_removal_success(self): result = await wait_for_ci_gate(state) # Verify remove_labels called with correct parameters - mock_jira.remove_labels.assert_called_once_with( - "TEST-300", - ["forge:implementing"] - ) + mock_jira.remove_labels.assert_called_once_with("TEST-300", ["forge:implementing"]) # Verify workflow continues assert result["is_paused"] is True assert result["current_node"] == "wait_for_ci_gate" @@ -213,8 +215,11 @@ async def test_label_removal_label_not_present(self, caplog): assert result["is_paused"] is True assert result["current_node"] == "wait_for_ci_gate" # Verify error logged (via post_status_comment utility) - assert any("Failed to remove implementing label" in record.message - for record in caplog.records if record.levelname == "WARNING") + assert any( + "Failed to remove implementing label" in record.message + for record in caplog.records + if record.levelname == "WARNING" + ) @pytest.mark.asyncio async def test_label_removal_api_error(self, caplog): @@ -238,8 +243,11 @@ async def test_label_removal_api_error(self, caplog): assert result["is_paused"] is True assert result["current_node"] == "wait_for_ci_gate" # Verify error logged at WARNING level - assert any("Failed to remove implementing label" in record.message - for record in caplog.records if record.levelname == "WARNING") + assert any( + "Failed to remove implementing label" in record.message + for record in caplog.records + if record.levelname == "WARNING" + ) @pytest.mark.asyncio async def test_label_removal_not_called_on_reentry(self): @@ -282,10 +290,8 @@ async def test_label_addition_success(self): # Verify set_workflow_label called with forge:ci-pending from forge.models.workflow import ForgeLabel - mock_jira.set_workflow_label.assert_called_once_with( - "TEST-400", - ForgeLabel.TASK_CI_PENDING - ) + + mock_jira.set_workflow_label.assert_called_once_with("TEST-400", ForgeLabel.TASK_CI_PENDING) # Verify workflow continues assert result["is_paused"] is True assert result["current_node"] == "wait_for_ci_gate" @@ -312,8 +318,11 @@ async def test_label_addition_api_error(self, caplog): assert result["is_paused"] is True assert result["current_node"] == "wait_for_ci_gate" # Verify error logged at WARNING level - assert any("Failed to set ci-pending label" in record.message - for record in caplog.records if record.levelname == "WARNING") + assert any( + "Failed to set ci-pending label" in record.message + for record in caplog.records + if record.levelname == "WARNING" + ) @pytest.mark.asyncio async def test_label_addition_not_called_on_reentry(self): @@ -359,8 +368,11 @@ async def test_comment_posting_error_logged_and_suppressed(self, caplog): assert result["is_paused"] is True assert result["current_node"] == "wait_for_ci_gate" # Verify error logged - assert any("Failed to post status comment" in record.message - for record in caplog.records if record.levelname == "WARNING") + assert any( + "Failed to post status comment" in record.message + for record in caplog.records + if record.levelname == "WARNING" + ) @pytest.mark.asyncio async def test_label_removal_error_logged_and_suppressed(self, caplog): @@ -382,8 +394,11 @@ async def test_label_removal_error_logged_and_suppressed(self, caplog): # Verify workflow continues assert result["is_paused"] is True # Verify error logged - assert any("Failed to remove implementing label" in record.message - for record in caplog.records if record.levelname == "WARNING") + assert any( + "Failed to remove implementing label" in record.message + for record in caplog.records + if record.levelname == "WARNING" + ) @pytest.mark.asyncio async def test_label_addition_error_logged_and_suppressed(self, caplog): @@ -405,8 +420,11 @@ async def test_label_addition_error_logged_and_suppressed(self, caplog): # Verify workflow continues assert result["is_paused"] is True # Verify error logged - assert any("Failed to set ci-pending label" in record.message - for record in caplog.records if record.levelname == "WARNING") + assert any( + "Failed to set ci-pending label" in record.message + for record in caplog.records + if record.levelname == "WARNING" + ) @pytest.mark.asyncio async def test_all_operations_fail_workflow_still_continues(self, caplog): @@ -432,7 +450,9 @@ async def test_all_operations_fail_workflow_still_continues(self, caplog): assert result["is_paused"] is True assert result["current_node"] == "wait_for_ci_gate" # Verify all errors logged - warning_messages = [record.message for record in caplog.records if record.levelname == "WARNING"] + warning_messages = [ + record.message for record in caplog.records if record.levelname == "WARNING" + ] assert any("Failed to post status comment" in msg for msg in warning_messages) assert any("Failed to remove implementing label" in msg for msg in warning_messages) assert any("Failed to set ci-pending label" in msg for msg in warning_messages) diff --git a/tests/unit/workflow/test_yolo_mode.py b/tests/unit/workflow/test_yolo_mode.py index b4a261c1..c6958ea0 100644 --- a/tests/unit/workflow/test_yolo_mode.py +++ b/tests/unit/workflow/test_yolo_mode.py @@ -2,9 +2,10 @@ import pytest -from forge.models.workflow import ForgeLabel, TicketType -from forge.workflow.feature.state import create_initial_feature_state +from forge.models.workflow import ForgeLabel +from forge.queue.models import QueueMessage from forge.workflow.bug.state import create_initial_bug_state +from forge.workflow.feature.state import create_initial_feature_state class TestForgeLabelYolo: @@ -38,7 +39,9 @@ class TestBuildInitialStateYoloMode: def _make_worker(self): from unittest.mock import MagicMock + from forge.orchestrator.worker import OrchestratorWorker + worker = OrchestratorWorker.__new__(OrchestratorWorker) worker.settings = MagicMock() worker.router = MagicMock() @@ -46,7 +49,9 @@ def _make_worker(self): def _make_message(self, labels: list): from unittest.mock import MagicMock + from forge.models.events import EventSource + msg = MagicMock() msg.ticket_key = "TEST-1" msg.source = EventSource.JIRA @@ -83,7 +88,9 @@ def test_yolo_mode_false_when_no_labels(self): def test_yolo_mode_false_for_github_source(self): from unittest.mock import MagicMock + from forge.models.events import EventSource + msg = MagicMock() msg.ticket_key = "TEST-1" msg.source = EventSource.GITHUB @@ -99,9 +106,12 @@ def test_yolo_mode_false_for_github_source(self): class TestYoloLabelAddedMidWorkflow: """When forge:yolo is added while paused at a gate, yolo_mode is set and workflow unpauses.""" - def _make_yolo_label_message(self, current_labels: str, previous_labels: str = "") -> "QueueMessage": + def _make_yolo_label_message( + self, current_labels: str, previous_labels: str = "" + ) -> "QueueMessage": from forge.models.events import EventSource from forge.queue.models import QueueMessage + return QueueMessage( message_id="1234567890-0", event_id="test-event-yolo", @@ -139,6 +149,7 @@ def _make_gate_state(self, current_node: str, **extra) -> dict: @pytest.mark.asyncio async def test_yolo_label_addition_at_prd_gate_activates_yolo(self): from forge.orchestrator.worker import OrchestratorWorker + worker = OrchestratorWorker(consumer_name="test-worker") message = self._make_yolo_label_message( current_labels="forge:managed forge:yolo", @@ -152,6 +163,7 @@ async def test_yolo_label_addition_at_prd_gate_activates_yolo(self): @pytest.mark.asyncio async def test_yolo_label_addition_outside_gate_does_not_activate(self): from forge.orchestrator.worker import OrchestratorWorker + worker = OrchestratorWorker(consumer_name="test-worker") message = self._make_yolo_label_message( current_labels="forge:managed forge:yolo", @@ -166,6 +178,7 @@ async def test_yolo_label_addition_outside_gate_does_not_activate(self): @pytest.mark.asyncio async def test_yolo_label_already_present_does_not_re_trigger(self): from forge.orchestrator.worker import OrchestratorWorker + worker = OrchestratorWorker(consumer_name="test-worker") # forge:yolo was already in fromString — not a new addition message = self._make_yolo_label_message( @@ -184,6 +197,7 @@ class TestYoloGateRouting: def _feature_state(self, current_node: str, **extra) -> dict: from forge.workflow.feature.state import create_initial_feature_state + state = create_initial_feature_state("TEST-1") state["current_node"] = current_node state["is_paused"] = True @@ -193,28 +207,36 @@ def _feature_state(self, current_node: str, **extra) -> dict: def test_prd_route_auto_approves_in_yolo_mode(self): from forge.workflow.gates.prd_approval import route_prd_approval + state = self._feature_state("prd_approval_gate", prd_content="# PRD") assert route_prd_approval(state) == "generate_spec" def test_spec_route_auto_approves_in_yolo_mode(self): from forge.workflow.gates.spec_approval import route_spec_approval + state = self._feature_state("spec_approval_gate", spec_content="# Spec") assert route_spec_approval(state) == "decompose_epics" - def test_plan_route_auto_approves_in_yolo_mode(self): + @pytest.mark.asyncio + async def test_plan_route_auto_approves_in_yolo_mode(self): from forge.workflow.gates.plan_approval import route_plan_approval + state = self._feature_state("plan_approval_gate", epic_keys=["EPIC-1"]) - assert route_plan_approval(state) == "generate_tasks" + assert await route_plan_approval(state) == "generate_tasks" - def test_task_route_auto_approves_in_yolo_mode(self): + @pytest.mark.asyncio + async def test_task_route_auto_approves_in_yolo_mode(self): from forge.workflow.gates.task_approval import route_task_approval + state = self._feature_state("task_approval_gate", task_keys=["TASK-1"]) - assert route_task_approval(state) == "task_router" + assert await route_task_approval(state) == "task_router" def test_yolo_false_still_pauses_at_prd_gate(self): from langgraph.graph import END - from forge.workflow.gates.prd_approval import route_prd_approval + from forge.workflow.feature.state import create_initial_feature_state + from forge.workflow.gates.prd_approval import route_prd_approval + state = create_initial_feature_state("TEST-1") state["current_node"] = "prd_approval_gate" state["is_paused"] = True @@ -224,6 +246,7 @@ def test_yolo_false_still_pauses_at_prd_gate(self): def test_yolo_does_not_override_question_routing(self): from forge.workflow.gates.prd_approval import route_prd_approval + state = self._feature_state("prd_approval_gate", prd_content="# PRD") state["is_question"] = True state["feedback_comment"] = "?Why REST?" @@ -259,6 +282,7 @@ def _rca_state(self, **extra) -> dict: @pytest.mark.asyncio async def test_yolo_selects_option_1_without_pausing(self): from unittest.mock import AsyncMock, patch + from forge.workflow.nodes.rca_option_gate import rca_option_gate state = self._rca_state() @@ -278,6 +302,7 @@ async def test_yolo_selects_option_1_without_pausing(self): async def test_yolo_still_posts_rca_comment(self): """RCA comment is posted even in yolo mode (audit trail preserved).""" from unittest.mock import AsyncMock, patch + from forge.workflow.nodes.rca_option_gate import rca_option_gate state = self._rca_state() @@ -295,6 +320,7 @@ async def test_yolo_still_posts_rca_comment(self): async def test_non_yolo_still_pauses(self): """With yolo_mode=False, gate pauses normally.""" from unittest.mock import AsyncMock, patch + from forge.workflow.nodes.rca_option_gate import rca_option_gate state = self._rca_state(yolo_mode=False) diff --git a/tests/unit/workflow/utils/test_draft_manager.py b/tests/unit/workflow/utils/test_draft_manager.py new file mode 100644 index 00000000..04c13f76 --- /dev/null +++ b/tests/unit/workflow/utils/test_draft_manager.py @@ -0,0 +1,307 @@ +# mypy: disallow-untyped-decorators=False +"""Tests for DraftManager utility class.""" + +from datetime import UTC, datetime +from unittest.mock import AsyncMock, MagicMock + +import pytest +from pytest import LogCaptureFixture + +from forge.integrations.jira import JiraClient +from forge.models.draft import DraftItem, ForgeDecompositionDraft +from forge.workflow.utils.draft_manager import ( + FORGE_STORIES_DRAFT_FILENAME, + FORGE_TASKS_DRAFT_FILENAME, + DraftManager, +) + + +@pytest.fixture( + params=[ + ("stories", FORGE_STORIES_DRAFT_FILENAME), + ("tasks", FORGE_TASKS_DRAFT_FILENAME), + ] +) +def draft_config(request: pytest.FixtureRequest) -> tuple[str, str]: + """Return a tuple of (phase, filename) representing draft configurations.""" + val: tuple[str, str] = request.param + return val + + +@pytest.fixture +def sample_draft(draft_config: tuple[str, str]) -> ForgeDecompositionDraft: + """Return a valid ForgeDecompositionDraft instance matching the draft configuration.""" + phase, _ = draft_config + now = datetime.now(UTC) + return ForgeDecompositionDraft( + parent_key="PROJ-123", + phase=phase, + items=[ + DraftItem( + id=1, + summary=f"{phase.capitalize()} 1", + description="Desc 1", + repo="repo-a", + acceptance_criteria=["AC 1"], + ) + ], + version=1, + created_at=now, + updated_at=now, + ) + + +class TestDraftManager: + """Test cases for DraftManager CRUD operations on Jira parent tickets.""" + + @pytest.mark.asyncio + async def test_save_draft_attachment_success_no_existing( + self, draft_config: tuple[str, str], sample_draft: ForgeDecompositionDraft + ) -> None: + """Should successfully upload draft when no prior matching attachment exists.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.delete_attachments_by_name = AsyncMock(return_value=0) + mock_jira.add_attachment = AsyncMock() + + await DraftManager.save_draft_attachment(mock_jira, "PROJ-123", sample_draft, filename) + + mock_jira.delete_attachments_by_name.assert_called_once_with("PROJ-123", filename) + + # Verify serialized content passed to add_attachment + expected_bytes = sample_draft.model_dump_json().encode("utf-8") + mock_jira.add_attachment.assert_called_once_with("PROJ-123", filename, expected_bytes) + + @pytest.mark.asyncio + async def test_save_draft_attachment_success_with_existing( + self, draft_config: tuple[str, str], sample_draft: ForgeDecompositionDraft + ) -> None: + """Should delete existing matching attachment before uploading new draft.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.delete_attachments_by_name = AsyncMock(return_value=2) + mock_jira.add_attachment = AsyncMock() + + await DraftManager.save_draft_attachment(mock_jira, "PROJ-123", sample_draft, filename) + + mock_jira.delete_attachments_by_name.assert_called_once_with("PROJ-123", filename) + + expected_bytes = sample_draft.model_dump_json().encode("utf-8") + mock_jira.add_attachment.assert_called_once_with("PROJ-123", filename, expected_bytes) + + @pytest.mark.asyncio + async def test_save_draft_attachment_failure_to_delete( + self, draft_config: tuple[str, str], sample_draft: ForgeDecompositionDraft + ) -> None: + """Should propagate delete_attachments_by_name exception.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.delete_attachments_by_name = AsyncMock(side_effect=Exception("Delete Error")) + + with pytest.raises(Exception, match="Delete Error"): + await DraftManager.save_draft_attachment(mock_jira, "PROJ-123", sample_draft, filename) + + @pytest.mark.asyncio + async def test_save_draft_attachment_failure_to_upload( + self, draft_config: tuple[str, str], sample_draft: ForgeDecompositionDraft + ) -> None: + """Should propagate add_attachment exception.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.delete_attachments_by_name = AsyncMock(return_value=0) + mock_jira.add_attachment = AsyncMock(side_effect=Exception("Upload Error")) + + with pytest.raises(Exception, match="Upload Error"): + await DraftManager.save_draft_attachment(mock_jira, "PROJ-123", sample_draft, filename) + + @pytest.mark.asyncio + async def test_get_draft_attachment_success( + self, draft_config: tuple[str, str], sample_draft: ForgeDecompositionDraft + ) -> None: + """Should download and successfully parse draft attachment if found.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.get_attachments = AsyncMock( + return_value=[ + { + "id": "att-222", + "filename": filename, + "content_url": "http://url2", + }, + ] + ) + serialized_bytes = sample_draft.model_dump_json().encode("utf-8") + mock_jira.download_attachment = AsyncMock(return_value=serialized_bytes) + + res = await DraftManager.get_draft_attachment(mock_jira, "PROJ-123", filename) + + assert res is not None + assert res.parent_key == sample_draft.parent_key + assert res.phase == sample_draft.phase + assert len(res.items) == len(sample_draft.items) + assert res.items[0].summary == sample_draft.items[0].summary + mock_jira.get_attachments.assert_called_once_with("PROJ-123") + mock_jira.download_attachment.assert_called_once_with("http://url2") + + @pytest.mark.asyncio + async def test_get_draft_attachment_success_alternate_url_key( + self, draft_config: tuple[str, str], sample_draft: ForgeDecompositionDraft + ) -> None: + """Should download using 'content_url' key.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.get_attachments = AsyncMock( + return_value=[ + { + "id": "att-222", + "filename": filename, + "content_url": "http://url2-alt", + }, + ] + ) + serialized_bytes = sample_draft.model_dump_json().encode("utf-8") + mock_jira.download_attachment = AsyncMock(return_value=serialized_bytes) + + res = await DraftManager.get_draft_attachment(mock_jira, "PROJ-123", filename) + + assert res is not None + mock_jira.download_attachment.assert_called_once_with("http://url2-alt") + + @pytest.mark.asyncio + async def test_get_draft_attachment_not_found(self, draft_config: tuple[str, str]) -> None: + """Should return None if attachment does not exist.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.get_attachments = AsyncMock(return_value=[]) + + res = await DraftManager.get_draft_attachment(mock_jira, "PROJ-123", filename) + + assert res is None + mock_jira.get_attachments.assert_called_once_with("PROJ-123") + + @pytest.mark.asyncio + async def test_get_draft_attachment_missing_content_url( + self, draft_config: tuple[str, str] + ) -> None: + """Should log warning and return None if content URL is missing.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.get_attachments = AsyncMock( + return_value=[ + {"id": "att-222", "filename": filename}, + ] + ) + + res = await DraftManager.get_draft_attachment(mock_jira, "PROJ-123", filename) + + assert res is None + + @pytest.mark.asyncio + async def test_get_draft_attachment_validation_failure( + self, draft_config: tuple[str, str], caplog: LogCaptureFixture + ) -> None: + """Should log a warning and return None if draft JSON is invalid according to model schema.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.get_attachments = AsyncMock( + return_value=[ + { + "id": "att-222", + "filename": filename, + "content_url": "http://url2", + }, + ] + ) + # Missing required fields like updated_at, parent_key, etc. + invalid_bytes = b'{"parent_key": "PROJ-123", "phase": "invalid_phase"}' + mock_jira.download_attachment = AsyncMock(return_value=invalid_bytes) + + res = await DraftManager.get_draft_attachment(mock_jira, "PROJ-123", filename) + + assert res is None + # Verify warning log was printed + assert any( + "Validation failed for draft attachment" in record.message + and record.levelname == "WARNING" + for record in caplog.records + ) + + @pytest.mark.asyncio + async def test_get_draft_attachment_parsing_failure( + self, draft_config: tuple[str, str], caplog: LogCaptureFixture + ) -> None: + """Should log a warning and return None if draft bytes are completely non-JSON.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.get_attachments = AsyncMock( + return_value=[ + { + "id": "att-222", + "filename": filename, + "content_url": "http://url2", + }, + ] + ) + mock_jira.download_attachment = AsyncMock(return_value=b"not json at all") + + res = await DraftManager.get_draft_attachment(mock_jira, "PROJ-123", filename) + + assert res is None + # Verify warning log was printed + assert any( + "Failed to parse draft attachment" in record.message and record.levelname == "WARNING" + for record in caplog.records + ) + + @pytest.mark.asyncio + async def test_get_draft_attachment_download_failure( + self, draft_config: tuple[str, str] + ) -> None: + """Should propagate download_attachment exceptions.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.get_attachments = AsyncMock( + return_value=[ + { + "id": "att-222", + "filename": filename, + "content_url": "http://url2", + }, + ] + ) + mock_jira.download_attachment = AsyncMock(side_effect=Exception("Network Timeout")) + + with pytest.raises(Exception, match="Network Timeout"): + await DraftManager.get_draft_attachment(mock_jira, "PROJ-123", filename) + + @pytest.mark.asyncio + async def test_delete_draft_attachment_success(self, draft_config: tuple[str, str]) -> None: + """Should delete all matching attachments if found.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.delete_attachments_by_name = AsyncMock(return_value=2) + + await DraftManager.delete_draft_attachment(mock_jira, "PROJ-123", filename) + + mock_jira.delete_attachments_by_name.assert_called_once_with("PROJ-123", filename) + + @pytest.mark.asyncio + async def test_delete_draft_attachment_not_found(self, draft_config: tuple[str, str]) -> None: + """Should do nothing and succeed if no matching attachment found.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.delete_attachments_by_name = AsyncMock(return_value=0) + + await DraftManager.delete_draft_attachment(mock_jira, "PROJ-123", filename) + + mock_jira.delete_attachments_by_name.assert_called_once_with("PROJ-123", filename) + + @pytest.mark.asyncio + async def test_delete_draft_attachment_failure(self, draft_config: tuple[str, str]) -> None: + """Should propagate delete_attachments_by_name exception.""" + _, filename = draft_config + mock_jira = MagicMock(spec=JiraClient) + mock_jira.delete_attachments_by_name = AsyncMock(side_effect=Exception("Delete Error")) + + with pytest.raises(Exception, match="Delete Error"): + await DraftManager.delete_draft_attachment(mock_jira, "PROJ-123", filename) diff --git a/tests/unit/workflow/utils/test_jira_status.py b/tests/unit/workflow/utils/test_jira_status.py index 644835cd..92c49782 100644 --- a/tests/unit/workflow/utils/test_jira_status.py +++ b/tests/unit/workflow/utils/test_jira_status.py @@ -140,18 +140,15 @@ async def test_transition_tasks_success(self, caplog) -> None: # Verify success logs for each task assert any( - "Transitioned TASK-1 to In Progress" in record.message - and record.levelname == "INFO" + "Transitioned TASK-1 to In Progress" in record.message and record.levelname == "INFO" for record in caplog.records ) assert any( - "Transitioned TASK-2 to In Progress" in record.message - and record.levelname == "INFO" + "Transitioned TASK-2 to In Progress" in record.message and record.levelname == "INFO" for record in caplog.records ) assert any( - "Transitioned TASK-3 to In Progress" in record.message - and record.levelname == "INFO" + "Transitioned TASK-3 to In Progress" in record.message and record.levelname == "INFO" for record in caplog.records ) @@ -175,13 +172,11 @@ async def transition_side_effect(task_key: str, _status: str): # Verify success logs for tasks 1 and 3 assert any( - "Transitioned TASK-1 to In Progress" in record.message - and record.levelname == "INFO" + "Transitioned TASK-1 to In Progress" in record.message and record.levelname == "INFO" for record in caplog.records ) assert any( - "Transitioned TASK-3 to In Progress" in record.message - and record.levelname == "INFO" + "Transitioned TASK-3 to In Progress" in record.message and record.levelname == "INFO" for record in caplog.records ) @@ -213,13 +208,11 @@ async def transition_side_effect(task_key: str, _status: str): # Verify success logs for tasks 1 and 3 assert any( - "Transitioned TASK-1 to In Progress" in record.message - and record.levelname == "INFO" + "Transitioned TASK-1 to In Progress" in record.message and record.levelname == "INFO" for record in caplog.records ) assert any( - "Transitioned TASK-3 to In Progress" in record.message - and record.levelname == "INFO" + "Transitioned TASK-3 to In Progress" in record.message and record.levelname == "INFO" for record in caplog.records ) diff --git a/tests/unit/workflow/utils/test_review_report.py b/tests/unit/workflow/utils/test_review_report.py index 248ec042..f492d583 100644 --- a/tests/unit/workflow/utils/test_review_report.py +++ b/tests/unit/workflow/utils/test_review_report.py @@ -346,6 +346,7 @@ def test_second_call_preserves_first_exhaustion_entry(self): prior entry. When a single node calls it twice (e.g., ci_evaluator for analyze_ci then fix_ci), the second call silently drops the first. """ + def _exhausted_result(task_key: str, step_name: str) -> ContainerResult: return ContainerResult( success=True, @@ -368,7 +369,9 @@ def _exhausted_result(task_key: str, step_name: str) -> ContainerResult: state: dict = {} # First call: analyze_ci step exhausted - state = merge_review_exhaustion(state, _exhausted_result("T-1", "analyze_ci"), "T-1", "analyze_ci") + state = merge_review_exhaustion( + state, _exhausted_result("T-1", "analyze_ci"), "T-1", "analyze_ci" + ) assert "T-1__analyze_ci" in state["review_exhaustion_report"] # Second call: fix_ci step also exhausted diff --git a/tests/unit/workspace/test_git_ops_redaction.py b/tests/unit/workspace/test_git_ops_redaction.py index 640e0705..6741eca6 100644 --- a/tests/unit/workspace/test_git_ops_redaction.py +++ b/tests/unit/workspace/test_git_ops_redaction.py @@ -52,9 +52,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/workflow/test_draft_review_flow.py b/tests/workflow/test_draft_review_flow.py new file mode 100644 index 00000000..d764fe20 --- /dev/null +++ b/tests/workflow/test_draft_review_flow.py @@ -0,0 +1,831 @@ +"""Integration tests for Draft Review Flow. + +Covers YOLO bypass path, draft attachment creation/cleanup, BR-003 truncation +rules, excluded item skipping during ticket provisioning, and draft retention +on partial ticket provisioning failure. +""" + +from datetime import UTC, datetime +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from forge.config import Settings +from forge.models.draft import DraftItem, ForgeDecompositionDraft +from forge.workflow.gates.plan_approval import provision_epics_from_draft, route_plan_approval +from forge.workflow.gates.task_approval import provision_tasks_from_draft, route_task_approval +from forge.workflow.nodes.epic_decomposition import decompose_epics +from forge.workflow.nodes.task_generation import generate_tasks +from forge.workflow.utils.draft_manager import DraftManager + + +@pytest.fixture +def mock_settings() -> Settings: + """Create settings for tests.""" + return Settings( + redis_url="redis://localhost:6379/0", + jira_base_url="https://test.atlassian.net", + jira_api_token="test-token", + jira_user_email="test@example.com", + jira_webhook_secret="test-webhook-secret", + github_token="test-github-token", + github_webhook_secret="test-github-webhook-secret", + llm_backend="anthropic", + llm_model="claude-sonnet-4-5-20250929", + anthropic_api_key="test-anthropic-key", + yolo_mode=False, + ) + + +@pytest.fixture +def base_epic_state() -> dict[str, Any]: + """Base state for epic decomposition.""" + return { + "ticket_key": "TEST-100", + "spec_content": "Build feature x.", + "qa_history": [], + "retry_count": 0, + "yolo_mode": False, + "epic_keys": [], + } + + +@pytest.fixture +def base_task_state() -> dict[str, Any]: + """Base state for task generation.""" + return { + "ticket_key": "TEST-100", + "spec_content": "Build feature x.", + "qa_history": [], + "retry_count": 0, + "yolo_mode": False, + "epic_keys": ["TEST-101"], + "task_keys": [], + "tasks_by_repo": {}, + } + + +@pytest.fixture +def mock_parent_issue() -> Any: + """Mock Jira parent issue.""" + issue = MagicMock() + issue.project_key = "TEST" + issue.summary = "Test Feature Summary" + issue.description = "Test Feature Description" + return issue + + +@pytest.fixture +def mock_epic_issue() -> Any: + """Mock Jira Epic issue.""" + issue = MagicMock() + issue.project_key = "TEST" + issue.summary = "Test Epic Summary" + issue.description = "Test Epic Plan Description" + return issue + + +@pytest.fixture +def mock_epics_data() -> list[dict[str, Any]]: + """Mock generated epics data from LLM agent.""" + return [ + {"summary": "Epic 1", "plan": "Plan for epic 1", "repo": "acme/repo1"}, + {"summary": "Epic 2", "plan": "Plan for epic 2", "repo": "acme/repo2"}, + ] + + +@pytest.fixture +def mock_tasks_data() -> list[dict[str, Any]]: + """Mock generated tasks data from LLM agent.""" + return [ + {"summary": "Task 1", "description": "Desc for task 1", "repo": "acme/repo1"}, + {"summary": "Task 2", "description": "Desc for task 2", "repo": "acme/repo2"}, + ] + + +class TestYoloBypassPath: + """Acceptance Criterion: Integration tests verify the YOLO bypass path.""" + + @pytest.mark.asyncio + async def test_epic_decomposition_yolo_bypass( + self, + base_epic_state: dict[str, Any], + mock_parent_issue: Any, + mock_epics_data: list[dict[str, Any]], + mock_settings: Settings, + ) -> None: + """Verify decompose_epics provisions immediately and does NOT save draft attachments when YOLO is active.""" + state = {**base_epic_state, "yolo_mode": True} + + with ( + patch( + "forge.workflow.nodes.epic_decomposition.get_settings", return_value=mock_settings + ), + patch("forge.workflow.nodes.epic_decomposition.JiraClient") as MockJira, + patch("forge.workflow.nodes.epic_decomposition.ForgeAgent") as MockAgent, + patch( + "forge.workflow.nodes.epic_decomposition.DraftManager", wraps=DraftManager + ) as MockDraftManager, + patch("forge.workflow.nodes.epic_decomposition.post_qa_summary_if_needed"), + patch("forge.workflow.nodes.epic_decomposition.post_status_comment"), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(return_value=mock_parent_issue) + mock_jira.get_labels = AsyncMock(return_value=["forge:managed"]) + mock_jira.get_project_repos = AsyncMock(return_value=["acme/repo1", "acme/repo2"]) + mock_jira.create_epic = AsyncMock(side_effect=["TEST-101", "TEST-102"]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + mock_agent.generate_epics = AsyncMock(return_value=mock_epics_data) + + MockDraftManager.save_draft_attachment = AsyncMock() + MockDraftManager.delete_draft_attachment = AsyncMock() + + result = await decompose_epics(state) + + # Verify immediate provisioning of Epics + assert mock_jira.create_epic.call_count == 2 + mock_jira.create_epic.assert_any_call( + project_key="TEST", + summary="Epic 1", + description="Plan for epic 1", + parent_key="TEST-100", + labels=["forge:managed", "forge:parent:TEST-100", "repo:acme/repo1"], + ) + + # Verify draft was NOT saved or cleaned up + MockDraftManager.save_draft_attachment.assert_not_called() + MockDraftManager.delete_draft_attachment.assert_not_called() + + # Verify workflow pauses state is not set, instead keys are returned and transitions + assert result["epic_keys"] == ["TEST-101", "TEST-102"] + assert result.get("is_paused") is not True + assert result["current_node"] == "plan_approval_gate" + + @pytest.mark.asyncio + async def test_task_generation_yolo_bypass( + self, + base_task_state: dict[str, Any], + mock_parent_issue: Any, + mock_epic_issue: Any, + mock_tasks_data: list[dict[str, Any]], + mock_settings: Settings, + ) -> None: + """Verify generate_tasks provisions immediately and does NOT save draft attachments when YOLO is active.""" + state = {**base_task_state, "yolo_mode": True} + + with ( + patch("forge.workflow.nodes.task_generation.get_settings", return_value=mock_settings), + patch("forge.workflow.nodes.task_generation.JiraClient") as MockJira, + patch("forge.workflow.nodes.task_generation.ForgeAgent") as MockAgent, + patch( + "forge.workflow.nodes.task_generation.DraftManager", wraps=DraftManager + ) as MockDraftManager, + patch("forge.workflow.nodes.task_generation.post_status_comment"), + patch( + "forge.workflow.nodes.task_generation._generate_tasks_for_epic", + new_callable=AsyncMock, + return_value=mock_tasks_data, + ), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(side_effect=[mock_parent_issue, mock_epic_issue]) + mock_jira.get_labels = AsyncMock(return_value=["forge:managed"]) + mock_jira.create_task = AsyncMock(side_effect=["TEST-110", "TEST-111"]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + + MockDraftManager.save_draft_attachment = AsyncMock() + MockDraftManager.delete_draft_attachment = AsyncMock() + + result = await generate_tasks(state) + + # Verify immediate provisioning of Tasks + assert mock_jira.create_task.call_count == 2 + mock_jira.create_task.assert_any_call( + project_key="TEST", + summary="Task 1", + description="Desc for task 1", + parent_key="TEST-101", + labels=["forge:managed", "forge:parent:TEST-100", "repo:acme/repo1"], + ) + + # Verify draft was NOT saved or cleaned up + MockDraftManager.save_draft_attachment.assert_not_called() + MockDraftManager.delete_draft_attachment.assert_not_called() + + # Verify result state + assert result["task_keys"] == ["TEST-110", "TEST-111"] + assert result["tasks_by_repo"] == {"acme/repo1": ["TEST-110"], "acme/repo2": ["TEST-111"]} + assert result.get("is_paused") is not True + assert result["current_node"] == "task_approval_gate" + + +class TestDraftAttachmentCreationAndCleanup: + """Acceptance Criterion: Integration tests verify draft attachment creation and cleanup.""" + + @pytest.mark.asyncio + async def test_epic_decomposition_draft_review_flow( + self, + base_epic_state: dict[str, Any], + mock_parent_issue: Any, + mock_epics_data: list[dict[str, Any]], + mock_settings: Settings, + ) -> None: + """Verify that in non-YOLO mode, decompose_epics cleans up old drafts, saves the new draft JSON, posts comments, and pauses.""" + state = {**base_epic_state, "yolo_mode": False} + + with ( + patch( + "forge.workflow.nodes.epic_decomposition.get_settings", return_value=mock_settings + ), + patch("forge.workflow.nodes.epic_decomposition.JiraClient") as MockJira, + patch("forge.workflow.nodes.epic_decomposition.ForgeAgent") as MockAgent, + patch( + "forge.workflow.nodes.epic_decomposition.DraftManager", wraps=DraftManager + ) as MockDraftManager, + patch("forge.workflow.nodes.epic_decomposition.post_qa_summary_if_needed"), + patch("forge.workflow.nodes.epic_decomposition.post_status_comment"), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(return_value=mock_parent_issue) + mock_jira.get_labels = AsyncMock(return_value=["forge:managed"]) + mock_jira.get_project_repos = AsyncMock(return_value=["acme/repo1", "acme/repo2"]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + mock_agent.generate_epics = AsyncMock(return_value=mock_epics_data) + + MockDraftManager.delete_draft_attachment = AsyncMock() + MockDraftManager.save_draft_attachment = AsyncMock() + + result = await decompose_epics(state) + + # 1. Verify cleanup of any old drafts + MockDraftManager.delete_draft_attachment.assert_called_once_with( + mock_jira, "TEST-100", "forge-stories-draft.json" + ) + + # 2. Verify draft attachment saving + MockDraftManager.save_draft_attachment.assert_called_once() + saved_draft = MockDraftManager.save_draft_attachment.call_args[0][2] + assert isinstance(saved_draft, ForgeDecompositionDraft) + assert saved_draft.parent_key == "TEST-100" + assert saved_draft.phase == "stories" + assert len(saved_draft.items) == 2 + assert saved_draft.items[0].summary == "Epic 1" + assert saved_draft.items[0].repo == "acme/repo1" + + # 3. Verify comments posted + assert mock_jira.add_comment.call_count == 1 + comment_body = mock_jira.add_comment.call_args[0][1] + assert "### 📋 Proposed Epics Draft" in comment_body + assert "Epic 1" in comment_body + assert "/forge approve" in comment_body + + # 4. Verify workflow state transitions and pauses + assert result["is_paused"] is True + assert result["current_node"] == "plan_approval_gate" + assert result["epic_keys"] == [] + + @pytest.mark.asyncio + async def test_task_generation_draft_review_flow( + self, + base_task_state: dict[str, Any], + mock_parent_issue: Any, + mock_epic_issue: Any, + mock_tasks_data: list[dict[str, Any]], + mock_settings: Settings, + ) -> None: + """Verify that in non-YOLO mode, generate_tasks cleans up old drafts, saves the new draft JSON, posts comments, and pauses.""" + state = {**base_task_state, "yolo_mode": False} + + with ( + patch("forge.workflow.nodes.task_generation.get_settings", return_value=mock_settings), + patch("forge.workflow.nodes.task_generation.JiraClient") as MockJira, + patch("forge.workflow.nodes.task_generation.ForgeAgent") as MockAgent, + patch( + "forge.workflow.nodes.task_generation.DraftManager", wraps=DraftManager + ) as MockDraftManager, + patch("forge.workflow.nodes.task_generation.post_status_comment"), + patch( + "forge.workflow.nodes.task_generation._generate_tasks_for_epic", + new_callable=AsyncMock, + return_value=mock_tasks_data, + ), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(side_effect=[mock_parent_issue, mock_epic_issue]) + mock_jira.get_labels = AsyncMock(return_value=["forge:managed"]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + + MockDraftManager.delete_draft_attachment = AsyncMock() + MockDraftManager.save_draft_attachment = AsyncMock() + + result = await generate_tasks(state) + + # 1. Verify cleanup of any old drafts + MockDraftManager.delete_draft_attachment.assert_called_once_with( + mock_jira, "TEST-100", "forge-tasks-draft.json" + ) + + # 2. Verify draft attachment saving + MockDraftManager.save_draft_attachment.assert_called_once() + saved_draft = MockDraftManager.save_draft_attachment.call_args[0][2] + assert isinstance(saved_draft, ForgeDecompositionDraft) + assert saved_draft.parent_key == "TEST-100" + assert saved_draft.phase == "tasks" + assert len(saved_draft.items) == 2 + assert saved_draft.items[0].summary == "Task 1" + assert saved_draft.items[0].repo == "acme/repo1" + + # 3. Verify comments posted + assert mock_jira.add_comment.call_count == 1 + comment_body = mock_jira.add_comment.call_args[0][1] + assert "### 📋 Proposed Tasks Draft" in comment_body + assert "Task 1" in comment_body + assert "/forge approve" in comment_body + + # 4. Verify workflow state transitions and pauses + assert result["is_paused"] is True + assert result["current_node"] == "task_approval_gate" + assert result["task_keys"] == [] + + +class TestTruncationFallbackBoundaries: + """Acceptance Criterion: Integration tests verify character length and item count truncation fallback boundaries.""" + + @pytest.mark.asyncio + async def test_item_count_truncation_boundary_epics( + self, base_epic_state: dict[str, Any], mock_parent_issue: Any, mock_settings: Settings + ) -> None: + """BR-003: Verify that when item count > 15, the review comment is formatted in the condensed table format.""" + state = {**base_epic_state, "yolo_mode": False} + + # Generate 16 epics + many_epics = [ + {"summary": f"Epic {i}", "plan": f"Plan {i}", "repo": f"acme/repo{i}"} + for i in range(1, 17) + ] + + with ( + patch( + "forge.workflow.nodes.epic_decomposition.get_settings", return_value=mock_settings + ), + patch("forge.workflow.nodes.epic_decomposition.JiraClient") as MockJira, + patch("forge.workflow.nodes.epic_decomposition.ForgeAgent") as MockAgent, + patch( + "forge.workflow.nodes.epic_decomposition.DraftManager", wraps=DraftManager + ) as MockDraftManager, + patch("forge.workflow.nodes.epic_decomposition.post_qa_summary_if_needed"), + patch("forge.workflow.nodes.epic_decomposition.post_status_comment"), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(return_value=mock_parent_issue) + mock_jira.get_labels = AsyncMock(return_value=["forge:managed"]) + mock_jira.get_project_repos = AsyncMock(return_value=["acme/repo1"]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + mock_agent.generate_epics = AsyncMock(return_value=many_epics) + + MockDraftManager.delete_draft_attachment = AsyncMock() + MockDraftManager.save_draft_attachment = AsyncMock() + + await decompose_epics(state) + + # Verify comment is condensed + comment_body = mock_jira.add_comment.call_args[0][1] + assert "### 📋 Proposed Epics Draft (Condensed)" in comment_body + assert "Warning" in comment_body + assert "exceeds character or size limits" in comment_body + assert "forge-stories-draft.json" in comment_body + # Detailed descriptions of items should not be present + assert "#### 1. Epic 1" not in comment_body + + @pytest.mark.asyncio + async def test_item_count_truncation_boundary_tasks( + self, + base_task_state: dict[str, Any], + mock_parent_issue: Any, + mock_epic_issue: Any, + mock_settings: Settings, + ) -> None: + """BR-003: Verify that when item count > 15, the review comment is formatted in the condensed table format for tasks.""" + state = {**base_task_state, "yolo_mode": False} + + # Generate 16 tasks + many_tasks = [ + {"summary": f"Task {i}", "description": f"Desc {i}", "repo": f"acme/repo{i}"} + for i in range(1, 17) + ] + + with ( + patch("forge.workflow.nodes.task_generation.get_settings", return_value=mock_settings), + patch("forge.workflow.nodes.task_generation.JiraClient") as MockJira, + patch("forge.workflow.nodes.task_generation.ForgeAgent") as MockAgent, + patch( + "forge.workflow.nodes.task_generation.DraftManager", wraps=DraftManager + ) as MockDraftManager, + patch("forge.workflow.nodes.task_generation.post_status_comment"), + patch( + "forge.workflow.nodes.task_generation._generate_tasks_for_epic", + new_callable=AsyncMock, + return_value=many_tasks, + ), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(side_effect=[mock_parent_issue, mock_epic_issue]) + mock_jira.get_labels = AsyncMock(return_value=["forge:managed"]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + + MockDraftManager.delete_draft_attachment = AsyncMock() + MockDraftManager.save_draft_attachment = AsyncMock() + + await generate_tasks(state) + + # Verify comment is condensed + comment_body = mock_jira.add_comment.call_args[0][1] + assert "### 📋 Proposed Tasks Draft (Condensed)" in comment_body + assert "Warning" in comment_body + assert "exceeds character or size limits" in comment_body + assert "forge-tasks-draft.json" in comment_body + assert "#### 1. Task 1" not in comment_body + + @pytest.mark.asyncio + async def test_character_length_truncation_boundary_epics( + self, base_epic_state: dict[str, Any], mock_parent_issue: Any, mock_settings: Settings + ) -> None: + """BR-003: Verify that when comment character length > 32,767 characters, it falls back to a condensed table.""" + state = {**base_epic_state, "yolo_mode": False} + + # Create 1 huge plan for an epic + long_plan = "A" * 33000 + epics_with_long_plan = [ + {"summary": "Epic 1", "plan": long_plan, "repo": "acme/repo1"}, + {"summary": "Epic 2", "plan": "Short plan", "repo": "acme/repo2"}, + ] + + with ( + patch( + "forge.workflow.nodes.epic_decomposition.get_settings", return_value=mock_settings + ), + patch("forge.workflow.nodes.epic_decomposition.JiraClient") as MockJira, + patch("forge.workflow.nodes.epic_decomposition.ForgeAgent") as MockAgent, + patch( + "forge.workflow.nodes.epic_decomposition.DraftManager", wraps=DraftManager + ) as MockDraftManager, + patch("forge.workflow.nodes.epic_decomposition.post_qa_summary_if_needed"), + patch("forge.workflow.nodes.epic_decomposition.post_status_comment"), + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(return_value=mock_parent_issue) + mock_jira.get_labels = AsyncMock(return_value=["forge:managed"]) + mock_jira.get_project_repos = AsyncMock(return_value=["acme/repo1", "acme/repo2"]) + mock_jira.set_workflow_label = AsyncMock() + mock_jira.add_comment = AsyncMock() + + mock_agent = AsyncMock() + MockAgent.return_value = mock_agent + mock_agent.generate_epics = AsyncMock(return_value=epics_with_long_plan) + + MockDraftManager.delete_draft_attachment = AsyncMock() + MockDraftManager.save_draft_attachment = AsyncMock() + + await decompose_epics(state) + + # Verify comment is condensed due to length limit + comment_body = mock_jira.add_comment.call_args[0][1] + assert "### 📋 Proposed Epics Draft (Condensed)" in comment_body + assert "Warning" in comment_body + assert "exceeds character or size limits" in comment_body + assert "forge-stories-draft.json" in comment_body + # Detailed descriptions of items should not be present + assert "#### 1. Epic 1" not in comment_body + + +class TestApprovalCommandAndSkippingExcludedItems: + """Acceptance Criterion: Integration tests verify that excluded: true items are skipped during provisioning.""" + + @pytest.mark.asyncio + async def test_epics_provisioning_skips_excluded_items( + self, mock_parent_issue: Any, mock_settings: Settings + ) -> None: + """Verify only non-excluded items are provisioned and the attachment is deleted upon success.""" + state = { + "ticket_key": "TEST-100", + "is_paused": False, + "epic_keys": [], + } + + # Create draft where item 2 is excluded + draft = ForgeDecompositionDraft( + parent_key="TEST-100", + phase="stories", + items=[ + DraftItem( + id=1, + summary="Epic 1", + description="Plan 1", + repo="acme/repo1", + excluded=False, + acceptance_criteria=[], + ), + DraftItem( + id=2, + summary="Epic 2", + description="Plan 2", + repo="acme/repo2", + excluded=True, + acceptance_criteria=[], + ), + DraftItem( + id=3, + summary="Epic 3", + description="Plan 3", + repo="acme/repo3", + excluded=False, + acceptance_criteria=[], + ), + ], + version=1, + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + + with ( + patch("forge.config.get_settings", return_value=mock_settings), + patch("forge.integrations.jira.client.JiraClient") as MockJira, + patch("forge.workflow.utils.draft_manager.DraftManager") as MockDraftManager, + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(return_value=mock_parent_issue) + mock_jira.create_epic = AsyncMock(side_effect=["EPIC-1", "EPIC-3"]) + mock_jira.close = AsyncMock() + + MockDraftManager.get_draft_attachment = AsyncMock(return_value=draft) + MockDraftManager.delete_draft_attachment = AsyncMock() + + result_keys = await provision_epics_from_draft(state, mock_jira) + + # Verify only Epic 1 and Epic 3 were created + assert mock_jira.create_epic.call_count == 2 + mock_jira.create_epic.assert_any_call( + project_key="TEST", + summary="Epic 1", + description="Plan 1", + parent_key="TEST-100", + labels=["forge:managed", "forge:parent:TEST-100", "repo:acme/repo1"], + ) + mock_jira.create_epic.assert_any_call( + project_key="TEST", + summary="Epic 3", + description="Plan 3", + parent_key="TEST-100", + labels=["forge:managed", "forge:parent:TEST-100", "repo:acme/repo3"], + ) + + # Verify Epic 2 was skipped + for call in mock_jira.create_epic.call_args_list: + assert "Epic 2" not in call[1]["summary"] + + # Verify draft was deleted after successful provisioning + MockDraftManager.delete_draft_attachment.assert_called_once_with( + mock_jira, "TEST-100", "forge-stories-draft.json" + ) + assert result_keys == ["EPIC-1", "EPIC-3"] + + @pytest.mark.asyncio + async def test_tasks_provisioning_skips_excluded_items( + self, mock_parent_issue: Any, mock_settings: Settings + ) -> None: + """Verify only non-excluded tasks are provisioned and the attachment is deleted upon success.""" + state = { + "ticket_key": "TEST-100", + "is_paused": False, + "task_keys": [], + "epic_keys": ["EPIC-10"], + } + + # Create draft where item 2 is excluded + draft = ForgeDecompositionDraft( + parent_key="TEST-100", + phase="tasks", + items=[ + DraftItem( + id=1, + summary="Task 1", + description="Desc 1", + repo="acme/repo1", + excluded=False, + epic_key="EPIC-10", + acceptance_criteria=[], + ), + DraftItem( + id=2, + summary="Task 2", + description="Desc 2", + repo="acme/repo2", + excluded=True, + epic_key="EPIC-10", + acceptance_criteria=[], + ), + DraftItem( + id=3, + summary="Task 3", + description="Desc 3", + repo="acme/repo3", + excluded=False, + epic_key="EPIC-10", + acceptance_criteria=[], + ), + ], + version=1, + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + + with ( + patch("forge.config.get_settings", return_value=mock_settings), + patch("forge.integrations.jira.client.JiraClient") as MockJira, + patch("forge.workflow.utils.draft_manager.DraftManager") as MockDraftManager, + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(return_value=mock_parent_issue) + mock_jira.create_task = AsyncMock(side_effect=["TASK-1", "TASK-3"]) + mock_jira.close = AsyncMock() + + MockDraftManager.get_draft_attachment = AsyncMock(return_value=draft) + MockDraftManager.delete_draft_attachment = AsyncMock() + + task_keys, tasks_by_repo = await provision_tasks_from_draft(state, mock_jira) + + # Verify only Task 1 and Task 3 were created + assert mock_jira.create_task.call_count == 2 + mock_jira.create_task.assert_any_call( + project_key="TEST", + summary="Task 1", + description="Desc 1", + parent_key="EPIC-10", + labels=["forge:managed", "forge:parent:TEST-100", "repo:acme/repo1"], + ) + + # Verify draft was deleted after successful provisioning + MockDraftManager.delete_draft_attachment.assert_called_once_with( + mock_jira, "TEST-100", "forge-tasks-draft.json" + ) + assert task_keys == ["TASK-1", "TASK-3"] + assert tasks_by_repo == {"acme/repo1": ["TASK-1"], "acme/repo3": ["TASK-3"]} + + +class TestDraftRetentionOnFailure: + """Acceptance Criterion: Integration tests verify draft retention on partial ticket provisioning failure.""" + + @pytest.mark.asyncio + async def test_epics_provisioning_retains_draft_on_failure( + self, mock_parent_issue: Any, mock_settings: Settings + ) -> None: + """Verify draft is retained (delete_draft_attachment is not called) if epic provisioning fails midway.""" + state = { + "ticket_key": "TEST-100", + "is_paused": False, + "epic_keys": [], + } + + # Create draft with 2 epics + draft = ForgeDecompositionDraft( + parent_key="TEST-100", + phase="stories", + items=[ + DraftItem( + id=1, + summary="Epic 1", + description="Plan 1", + repo="acme/repo1", + excluded=False, + acceptance_criteria=[], + ), + DraftItem( + id=2, + summary="Epic 2", + description="Plan 2", + repo="acme/repo2", + excluded=False, + acceptance_criteria=[], + ), + ], + version=1, + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + + with ( + patch("forge.config.get_settings", return_value=mock_settings), + patch("forge.integrations.jira.client.JiraClient") as MockJira, + patch("forge.workflow.utils.draft_manager.DraftManager") as MockDraftManager, + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(return_value=mock_parent_issue) + # Epic 1 succeeds, Epic 2 fails with API error + mock_jira.create_epic = AsyncMock(side_effect=["EPIC-1", Exception("Jira API Failure")]) + mock_jira.close = AsyncMock() + + MockDraftManager.get_draft_attachment = AsyncMock(return_value=draft) + MockDraftManager.delete_draft_attachment = AsyncMock() + + with pytest.raises(Exception, match="Jira API Failure"): + await route_plan_approval(state) + + # Verify delete_draft_attachment was NEVER called, thus retaining the draft + MockDraftManager.delete_draft_attachment.assert_not_called() + + @pytest.mark.asyncio + async def test_tasks_provisioning_retains_draft_on_failure( + self, mock_parent_issue: Any, mock_settings: Settings + ) -> None: + """Verify draft is retained (delete_draft_attachment is not called) if task provisioning fails midway.""" + state = { + "ticket_key": "TEST-100", + "is_paused": False, + "task_keys": [], + "epic_keys": ["EPIC-10"], + } + + # Create draft with 2 tasks + draft = ForgeDecompositionDraft( + parent_key="TEST-100", + phase="tasks", + items=[ + DraftItem( + id=1, + summary="Task 1", + description="Desc 1", + repo="acme/repo1", + excluded=False, + epic_key="EPIC-10", + acceptance_criteria=[], + ), + DraftItem( + id=2, + summary="Task 2", + description="Desc 2", + repo="acme/repo2", + excluded=False, + epic_key="EPIC-10", + acceptance_criteria=[], + ), + ], + version=1, + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + + with ( + patch("forge.config.get_settings", return_value=mock_settings), + patch("forge.integrations.jira.client.JiraClient") as MockJira, + patch("forge.workflow.utils.draft_manager.DraftManager") as MockDraftManager, + ): + mock_jira = AsyncMock() + MockJira.return_value = mock_jira + mock_jira.get_issue = AsyncMock(return_value=mock_parent_issue) + # Task 1 succeeds, Task 2 fails with API error + mock_jira.create_task = AsyncMock(side_effect=["TASK-1", Exception("Jira API Failure")]) + mock_jira.close = AsyncMock() + + MockDraftManager.get_draft_attachment = AsyncMock(return_value=draft) + MockDraftManager.delete_draft_attachment = AsyncMock() + + with pytest.raises(Exception, match="Jira API Failure"): + await route_task_approval(state) + + # Verify delete_draft_attachment was NEVER called, thus retaining the draft + MockDraftManager.delete_draft_attachment.assert_not_called() diff --git a/tests/workflow/test_qualitative_review.py b/tests/workflow/test_qualitative_review.py index aa7c0e44..0e0cc329 100644 --- a/tests/workflow/test_qualitative_review.py +++ b/tests/workflow/test_qualitative_review.py @@ -151,7 +151,10 @@ async def test_run_qualitative_review_success_state_updates( with ( patch("forge.workflow.nodes.task_takeover_review.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.task_takeover_review.GitOperations") as mock_git, - patch("forge.workflow.nodes.task_takeover_review.ContainerRunner", return_value=mock_runner), + patch( + "forge.workflow.nodes.task_takeover_review.ContainerRunner", + return_value=mock_runner, + ), ): mock_git_instance = MagicMock() mock_git_instance._run_git = MagicMock() @@ -188,7 +191,10 @@ async def test_run_qualitative_review_failure_state_updates( with ( patch("forge.workflow.nodes.task_takeover_review.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.task_takeover_review.GitOperations") as mock_git, - patch("forge.workflow.nodes.task_takeover_review.ContainerRunner", return_value=mock_runner), + patch( + "forge.workflow.nodes.task_takeover_review.ContainerRunner", + return_value=mock_runner, + ), ): mock_git_instance = MagicMock() mock_git_instance._run_git = MagicMock() @@ -222,7 +228,10 @@ async def test_run_qualitative_review_retry_increment( with ( patch("forge.workflow.nodes.task_takeover_review.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.task_takeover_review.GitOperations") as mock_git, - patch("forge.workflow.nodes.task_takeover_review.ContainerRunner", return_value=mock_runner), + patch( + "forge.workflow.nodes.task_takeover_review.ContainerRunner", + return_value=mock_runner, + ), ): mock_git_instance = MagicMock() mock_git_instance._run_git = MagicMock() @@ -272,7 +281,10 @@ async def test_run_qualitative_review_valid_diff( with ( patch("forge.workflow.nodes.task_takeover_review.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.task_takeover_review.GitOperations") as mock_git, - patch("forge.workflow.nodes.task_takeover_review.ContainerRunner", return_value=mock_runner), + patch( + "forge.workflow.nodes.task_takeover_review.ContainerRunner", + return_value=mock_runner, + ), ): mock_git_instance = MagicMock() mock_git_instance._run_git = MagicMock() @@ -311,7 +323,10 @@ async def test_run_qualitative_review_invalid_diff_missing_tests( with ( patch("forge.workflow.nodes.task_takeover_review.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.task_takeover_review.GitOperations") as mock_git, - patch("forge.workflow.nodes.task_takeover_review.ContainerRunner", return_value=mock_runner), + patch( + "forge.workflow.nodes.task_takeover_review.ContainerRunner", + return_value=mock_runner, + ), ): mock_git_instance = MagicMock() mock_git_instance._run_git = MagicMock() @@ -351,7 +366,10 @@ async def test_run_qualitative_review_invalid_diff_unmet_criteria( with ( patch("forge.workflow.nodes.task_takeover_review.JiraClient", return_value=mock_jira), patch("forge.workflow.nodes.task_takeover_review.GitOperations") as mock_git, - patch("forge.workflow.nodes.task_takeover_review.ContainerRunner", return_value=mock_runner), + patch( + "forge.workflow.nodes.task_takeover_review.ContainerRunner", + return_value=mock_runner, + ), ): mock_git_instance = MagicMock() mock_git_instance._run_git = MagicMock() @@ -384,11 +402,12 @@ async def test_run_qualitative_review_exception_handling( mock_jira = _make_mock_jira() mock_jira.get_issue = AsyncMock(side_effect=RuntimeError("Jira API timeout")) - with patch( - "forge.workflow.nodes.task_takeover_review.JiraClient", return_value=mock_jira - ), patch( - "forge.workflow.nodes.task_takeover_review.prepare_workspace", - return_value=("/tmp/fake-workspace-review", MagicMock()), + with ( + patch("forge.workflow.nodes.task_takeover_review.JiraClient", return_value=mock_jira), + patch( + "forge.workflow.nodes.task_takeover_review.prepare_workspace", + return_value=("/tmp/fake-workspace-review", MagicMock()), + ), ): result = await run_qualitative_review(base_task_state) diff --git a/tests/workflow/utils/test_comment_command.py b/tests/workflow/utils/test_comment_command.py new file mode 100644 index 00000000..a833edf2 --- /dev/null +++ b/tests/workflow/utils/test_comment_command.py @@ -0,0 +1,475 @@ +"""Tests for parse_comment_command functionality.""" + +from typing import Any + +import pytest + +from forge.workflow.utils import parse_comment_command + + +def test_parse_remove_command_success() -> None: + """Test successful parsing of remove command.""" + result = parse_comment_command("/forge remove 2") + assert result == {"command": "remove", "id": 2} + + # Case insensitivity + result = parse_comment_command(" /FORGE Remove 42 ") + assert result == {"command": "remove", "id": 42} + + +def test_parse_remove_command_failures() -> None: + """Test parsing failures of remove command.""" + # Missing ID + result = parse_comment_command("/forge remove") + assert result is not None + assert "error" in result + assert result["command"] == "remove" + + # Invalid ID (string) + result = parse_comment_command("/forge remove abc") + assert result is not None + assert "error" in result + assert result["command"] == "remove" + + # Invalid ID (negative) + result = parse_comment_command("/forge remove -5") + assert result is not None + assert "error" in result + assert result["command"] == "remove" + + +def test_parse_exclude_command_success() -> None: + """Test successful parsing of exclude command.""" + result = parse_comment_command("/forge exclude 3") + assert result == {"command": "exclude", "id": 3} + + +def test_parse_exclude_command_failures() -> None: + """Test parsing failures of exclude command.""" + result = parse_comment_command("/forge exclude") + assert result is not None + assert "error" in result + assert result["command"] == "exclude" + + result = parse_comment_command("/forge exclude xyz") + assert result is not None + assert "error" in result + assert result["command"] == "exclude" + + +def test_parse_approve_command_success() -> None: + """Test successful parsing of approve command.""" + result = parse_comment_command("/forge approve") + assert result == {"command": "approve"} + + result = parse_comment_command(" /FORGE approve ") + assert result == {"command": "approve"} + + +def test_parse_approve_command_failures() -> None: + """Test parsing failures of approve command.""" + result = parse_comment_command("/forge approve 1") + assert result is not None + assert "error" in result + assert result["command"] == "approve" + + +def test_parse_add_command_success() -> None: + """Test successful parsing of add command.""" + result = parse_comment_command( + '/forge add summary="Implement API" repo="core-api" description="Set up endpoints"' + ) + assert result == { + "command": "add", + "params": { + "summary": "Implement API", + "repo": "core-api", + "description": "Set up endpoints", + }, + } + + # Mix of double, single and no quotes + result = parse_comment_command("/forge add summary='test single' count=42 name=\"quoted\"") + assert result == { + "command": "add", + "params": { + "summary": "test single", + "count": "42", + "name": "quoted", + }, + } + + +def test_parse_add_command_failures() -> None: + """Test parsing failures of add command.""" + # Missing parameters + result = parse_comment_command("/forge add") + assert result is not None + assert "error" in result + assert result["command"] == "add" + + # Malformed parameter (no key) + result = parse_comment_command("/forge add =value") + assert result is not None + assert "error" in result + assert result["command"] == "add" + + # Malformed parameters (trailing junk) + result = parse_comment_command('/forge add key="value" junk') + assert result is not None + assert "error" in result + assert result["command"] == "add" + + +def test_parse_update_command_success() -> None: + """Test successful parsing of update command.""" + result = parse_comment_command('/forge update 1 summary="New Summary"') + assert result == { + "command": "update", + "id": 1, + "params": {"summary": "New Summary"}, + } + + result = parse_comment_command("/forge update 100") + assert result == { + "command": "update", + "id": 100, + "params": {}, + } + + +def test_parse_update_command_failures() -> None: + """Test parsing failures of update command.""" + # Missing everything + result = parse_comment_command("/forge update") + assert result is not None + assert "error" in result + assert result["command"] == "update" + + # Missing ID but has parameters + result = parse_comment_command('/forge update summary="test"') + assert result is not None + assert "error" in result + assert result["command"] == "update" + + # Invalid ID + result = parse_comment_command('/forge update abc summary="test"') + assert result is not None + assert "error" in result + assert result["command"] == "update" + + # Malformed parameters + result = parse_comment_command('/forge update 1 summary="test" junk') + assert result is not None + assert "error" in result + assert result["command"] == "update" + + +def test_parse_command_non_matching() -> None: + """Test that unrelated texts or other /forge commands return None.""" + assert parse_comment_command("/forge skip-gate build") is None + assert parse_comment_command("/forge unskip-gate test") is None + assert parse_comment_command("/forge rebase") is None + assert parse_comment_command("/forge foo") is None + assert parse_comment_command("?what is this?") is None + assert parse_comment_command("!please update") is None + assert parse_comment_command("") is None + + +@pytest.fixture +def sample_draft_json() -> list[dict[str, Any]]: + return [ + { + "id": 1, + "summary": "Implement login route", + "description": "Create a POST route for user login", + "repo": "auth-api", + "acceptance_criteria": ["POST /login returns JWT on success"], + "excluded": False, + }, + { + "id": 2, + "summary": "Implement signup route", + "description": "Create a POST route for user registration", + "repo": "auth-api", + "acceptance_criteria": ["POST /signup registers user"], + "excluded": False, + }, + { + "id": 3, + "summary": "Add database migration", + "description": "Write migration script for users table", + "repo": "db-migration", + "acceptance_criteria": ["Users table has id, email, password"], + "excluded": True, + }, + ] + + +def test_apply_draft_modification_remove_success(sample_draft_json) -> None: + """Test successful removal and re-sequencing.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = {"command": "remove", "id": 2} + result = DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + assert len(result) == 2 + # Verify remaining items are re-sequenced + assert result[0]["id"] == 1 + assert result[0]["summary"] == "Implement login route" + assert result[1]["id"] == 2 + assert result[1]["summary"] == "Add database migration" + + +def test_apply_draft_modification_remove_missing_id(sample_draft_json) -> None: + """Test that removal fails if ID is missing.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = {"command": "remove"} + with pytest.raises(ValueError, match="Missing ID for removal"): + DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + +def test_apply_draft_modification_remove_not_found(sample_draft_json) -> None: + """Test that removal fails if ID is not found.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = {"command": "remove", "id": 99} + with pytest.raises(ValueError, match="Item with ID 99 not found for removal"): + DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + +def test_apply_draft_modification_add_success(sample_draft_json) -> None: + """Test successful addition with next sequential ID.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = { + "command": "add", + "params": { + "summary": "New task", + "description": "Task description", + "repo": "test-repo", + "acceptance_criteria": ["Criteria 1", "Criteria 2"], + }, + } + result = DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + assert len(result) == 4 + new_item = result[-1] + assert new_item["id"] == 4 + assert new_item["summary"] == "New task" + assert new_item["description"] == "Task description" + assert new_item["repo"] == "test-repo" + assert new_item["acceptance_criteria"] == ["Criteria 1", "Criteria 2"] + assert new_item["excluded"] is False + + +def test_apply_draft_modification_add_defaults(sample_draft_json) -> None: + """Test addition using only some parameters, relying on defaults for others.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = { + "command": "add", + "params": { + "summary": "Minimal task", + }, + } + result = DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + assert len(result) == 4 + new_item = result[-1] + assert new_item["id"] == 4 + assert new_item["summary"] == "Minimal task" + assert new_item["description"] == "" + assert new_item["repo"] == "" + assert new_item["acceptance_criteria"] == [] + assert new_item["excluded"] is False + + +def test_apply_draft_modification_update_success(sample_draft_json) -> None: + """Test successful update of target fields.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = { + "command": "update", + "id": 2, + "params": { + "summary": "Updated summary", + "acceptance_criteria": ["New AC"], + "excluded": True, + }, + } + result = DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + assert len(result) == 3 + updated_item = result[1] + assert updated_item["id"] == 2 + assert updated_item["summary"] == "Updated summary" + # Unchanged fields remain + assert updated_item["description"] == "Create a POST route for user registration" + assert updated_item["repo"] == "auth-api" + assert updated_item["acceptance_criteria"] == ["New AC"] + assert updated_item["excluded"] is True + + +def test_apply_draft_modification_update_missing_id(sample_draft_json) -> None: + """Test update raises error if ID is missing.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = {"command": "update", "params": {"summary": "No ID"}} + with pytest.raises(ValueError, match="Missing ID for update"): + DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + +def test_apply_draft_modification_update_not_found(sample_draft_json) -> None: + """Test update raises error if ID is not found.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = {"command": "update", "id": 99, "params": {"summary": "Not found"}} + with pytest.raises(ValueError, match="Item with ID 99 not found for update"): + DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + +def test_apply_draft_modification_exclude_success(sample_draft_json) -> None: + """Test flipping the excluded boolean key.""" + from forge.workflow.utils.draft_manager import DraftManager + + # Flip from False to True + parsed_cmd1 = {"command": "exclude", "id": 1} + result1 = DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd1) + assert result1[0]["excluded"] is True + + # Flip from True to False + parsed_cmd2 = {"command": "exclude", "id": 3} + result2 = DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd2) + assert result2[2]["excluded"] is False + + # Flip when excluded field is completely missing (defaults to False, so flips to True) + draft_without_excluded = [ + { + "id": 1, + "summary": "No excluded field", + "description": "Desc", + "repo": "repo", + "acceptance_criteria": [], + } + ] + parsed_cmd3 = {"command": "exclude", "id": 1} + result3 = DraftManager.apply_draft_modification(draft_without_excluded, parsed_cmd3) + assert result3[0]["excluded"] is True + + +def test_apply_draft_modification_exclude_missing_id(sample_draft_json) -> None: + """Test exclude raises error if ID is missing.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = {"command": "exclude"} + with pytest.raises(ValueError, match="Missing ID for exclude command"): + DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + +def test_apply_draft_modification_exclude_not_found(sample_draft_json) -> None: + """Test exclude raises error if ID is not found.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = {"command": "exclude", "id": 99} + with pytest.raises(ValueError, match="Item with ID 99 not found for exclude"): + DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + +@pytest.mark.parametrize( + "invalid_params, expected_error", + [ + ({"summary": 123}, "Field 'summary' must be a string"), + ({"description": ["not a string"]}, "Field 'description' must be a string"), + ({"repo": True}, "Field 'repo' must be a string"), + ( + {"acceptance_criteria": "string-instead-of-list"}, + "Field 'acceptance_criteria' must be a list of strings", + ), + ({"acceptance_criteria": [123]}, "Field 'acceptance_criteria' must be a list of strings"), + ({"excluded": "True"}, "Field 'excluded' must be a boolean"), + ({"unknown_field": "some-val"}, "Unknown field 'unknown_field'"), + ], +) +def test_apply_draft_modification_type_validation_add( + sample_draft_json, invalid_params, expected_error +) -> None: + """Test strict type validation for the 'add' command.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = {"command": "add", "params": invalid_params} + with pytest.raises(ValueError, match=expected_error): + DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + +@pytest.mark.parametrize( + "invalid_params, expected_error", + [ + ({"summary": 123}, "Field 'summary' must be a string"), + ({"description": ["not a string"]}, "Field 'description' must be a string"), + ({"repo": True}, "Field 'repo' must be a string"), + ( + {"acceptance_criteria": "string-instead-of-list"}, + "Field 'acceptance_criteria' must be a list of strings", + ), + ({"acceptance_criteria": [123]}, "Field 'acceptance_criteria' must be a list of strings"), + ({"excluded": "True"}, "Field 'excluded' must be a boolean"), + ({"unknown_field": "some-val"}, "Unknown field 'unknown_field'"), + ], +) +def test_apply_draft_modification_type_validation_update( + sample_draft_json, invalid_params, expected_error +) -> None: + """Test strict type validation for the 'update' command.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = {"command": "update", "id": 2, "params": invalid_params} + with pytest.raises(ValueError, match=expected_error): + DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + +def test_apply_draft_modification_parsing_error(sample_draft_json) -> None: + """Test that if the parsed_command dictionary contains an 'error' key, ValueError is raised.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = { + "command": "remove", + "error": "Missing integer ID for remove command", + } + with pytest.raises( + ValueError, match="Invalid command parameters: Missing integer ID for remove command" + ): + DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + +def test_apply_draft_modification_deepcopy_isolation(sample_draft_json) -> None: + """Test that apply_draft_modification doesn't modify the input draft_json in place.""" + from forge.workflow.utils.draft_manager import DraftManager + + parsed_cmd = { + "command": "update", + "id": 1, + "params": { + "summary": "Completely new summary", + }, + } + import copy + + original_copy = copy.deepcopy(sample_draft_json) + + result = DraftManager.apply_draft_modification(sample_draft_json, parsed_cmd) + + assert result[0]["summary"] == "Completely new summary" + assert sample_draft_json == original_copy + + +def test_apply_draft_modification_invalid_command_failures(sample_draft_json) -> None: + """Test that invalid command types raise ValueError.""" + from forge.workflow.utils.draft_manager import DraftManager + + with pytest.raises(ValueError, match="Command type is missing in parsed command."): + DraftManager.apply_draft_modification(sample_draft_json, {}) + + with pytest.raises(ValueError, match="Unsupported modification command type: 'invalid'"): + DraftManager.apply_draft_modification(sample_draft_json, {"command": "invalid"})