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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 4 additions & 3 deletions src/forge/integrations/agents/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -1162,7 +1162,9 @@ async def regenerate_document_with_feedback(
) -> ArtifactDocument:
"""Regenerate a PRD or specification with explicit repository selection."""
if content_type not in {"prd", "spec"}:
raise ValueError(f"Unsupported document type for structured regeneration: {content_type}")
raise ValueError(
f"Unsupported document type for structured regeneration: {content_type}"
)
prompt = load_prompt(
"regenerate",
content_type=content_type.upper(),
Expand Down Expand Up @@ -1254,8 +1256,7 @@ async def generate_epics(
)

epics = [
{"summary": epic.summary, "plan": epic.plan, "repo": epic.repo}
for epic in result.epics
{"summary": epic.summary, "plan": epic.plan, "repo": epic.repo} for epic in result.epics
]
logger.info(f"Generated {len(epics)} Epics")
return epics
Expand Down
4 changes: 3 additions & 1 deletion src/forge/workflow/nodes/task_generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -215,7 +215,9 @@ async def generate_tasks(state: WorkflowState) -> WorkflowState:
)
except Exception as exc:
logger.warning(
"Failed to report missing repository on Task %s: %s", task_key, exc
"Failed to report missing repository on Task %s: %s",
task_key,
exc,
)

# Assign the model tier for the newly created Task (BR-011).
Expand Down
71 changes: 53 additions & 18 deletions tests/flows/status_transitions/test_prd_rejected.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,12 +4,26 @@

import pytest

from forge.integrations.agents.structured_outputs import ArtifactDocument
from forge.models.workflow import TicketType
from forge.workflow.feature.state import create_initial_feature_state as create_initial_state
from forge.workflow.gates import route_prd_approval
from forge.workflow.nodes import regenerate_prd_with_feedback


@pytest.fixture(autouse=True)
def mock_repo_resolution(monkeypatch):
monkeypatch.setattr(
"forge.workflow.nodes.prd_generation.fetch_and_inject_references",
AsyncMock(side_effect=lambda _state, _jira, content: content),
)
monkeypatch.setattr(
"forge.workflow.nodes.prd_generation.get_effective_repos",
AsyncMock(return_value=["acme/repo"]),
)
monkeypatch.setattr("forge.workflow.nodes.prd_generation.reconcile_repo_labels", AsyncMock())


class TestPrdRejectedOnce:
"""Tests for single PRD rejection cycle."""

Expand Down Expand Up @@ -59,8 +73,9 @@ async def test_regeneration_incorporates_feedback(self, prd_pending_state):
mock_jira.close = AsyncMock()

mock_agent = MagicMock()
mock_agent.regenerate_with_feedback = AsyncMock(
return_value="""# Product Requirements Document
mock_agent.regenerate_document_with_feedback = AsyncMock(
return_value=ArtifactDocument(
content="""# Product Requirements Document

## Overview
Revised PRD with user personas.
Expand All @@ -72,7 +87,9 @@ async def test_regeneration_incorporates_feedback(self, prd_pending_state):
## Goals
- Enable user login
- Secure authentication
"""
""",
repositories=["acme/repo"],
)
)
mock_agent.close = AsyncMock()

Expand All @@ -83,14 +100,15 @@ async def test_regeneration_incorporates_feedback(self, prd_pending_state):
new_callable=AsyncMock,
return_value=prd_pending_state["prd_content"],
),
patch("forge.workflow.nodes.prd_generation.ensure_repo_labels", new_callable=AsyncMock),
patch("forge.workflow.stations.artifact_generation.ForgeAgent", return_value=mock_agent),
patch(
"forge.workflow.stations.artifact_generation.ForgeAgent", return_value=mock_agent
),
):
result = await regenerate_prd_with_feedback(prd_pending_state)

# Verify agent was called with feedback
mock_agent.regenerate_with_feedback.assert_called_once()
call_kwargs = mock_agent.regenerate_with_feedback.call_args.kwargs
mock_agent.regenerate_document_with_feedback.assert_called_once()
call_kwargs = mock_agent.regenerate_document_with_feedback.call_args.kwargs
assert "user persona" in call_kwargs["feedback"].lower()

# Verify new content
Expand All @@ -109,7 +127,9 @@ async def test_after_regeneration_returns_to_pending(self, prd_pending_state):
mock_jira.close = AsyncMock()

mock_agent = MagicMock()
mock_agent.regenerate_with_feedback = AsyncMock(return_value="# Revised PRD")
mock_agent.regenerate_document_with_feedback = AsyncMock(
return_value=ArtifactDocument(content="# Revised PRD", repositories=["acme/repo"])
)
mock_agent.close = AsyncMock()

with (
Expand All @@ -119,8 +139,9 @@ async def test_after_regeneration_returns_to_pending(self, prd_pending_state):
new_callable=AsyncMock,
return_value=prd_pending_state["prd_content"],
),
patch("forge.workflow.nodes.prd_generation.ensure_repo_labels", new_callable=AsyncMock),
patch("forge.workflow.stations.artifact_generation.ForgeAgent", return_value=mock_agent),
patch(
"forge.workflow.stations.artifact_generation.ForgeAgent", return_value=mock_agent
),
):
result = await regenerate_prd_with_feedback(prd_pending_state)

Expand Down Expand Up @@ -184,11 +205,15 @@ async def test_revision_count_increments(self, prd_state_first_revision):

mock_agent = MagicMock()
# Simulate error to increment retry count
mock_agent.regenerate_with_feedback = AsyncMock(side_effect=Exception("Simulated error"))
mock_agent.regenerate_document_with_feedback = AsyncMock(
side_effect=Exception("Simulated error")
)
mock_agent.close = AsyncMock()

with patch("forge.workflow.nodes.prd_generation.JiraClient", return_value=mock_jira):
with patch("forge.workflow.stations.artifact_generation.ForgeAgent", return_value=mock_agent):
with patch(
"forge.workflow.stations.artifact_generation.ForgeAgent", return_value=mock_agent
):
result = await regenerate_prd_with_feedback(prd_state_first_revision)

# Error case increments retry count
Expand Down Expand Up @@ -222,17 +247,22 @@ async def test_regeneration_uses_original_prd(self, prd_with_context):
mock_jira.update_description = AsyncMock()
mock_jira.add_comment = AsyncMock()
mock_jira.add_structured_comment = AsyncMock()
mock_jira.get_issue = AsyncMock(return_value=MagicMock(project_key="TEST"))
mock_jira.close = AsyncMock()

mock_agent = MagicMock()
mock_agent.regenerate_with_feedback = AsyncMock(return_value="# Revised")
mock_agent.regenerate_document_with_feedback = AsyncMock(
return_value=ArtifactDocument(content="# Revised", repositories=["acme/repo"])
)
mock_agent.close = AsyncMock()

with patch("forge.workflow.nodes.prd_generation.JiraClient", return_value=mock_jira):
with patch("forge.workflow.stations.artifact_generation.ForgeAgent", return_value=mock_agent):
with patch(
"forge.workflow.stations.artifact_generation.ForgeAgent", return_value=mock_agent
):
await regenerate_prd_with_feedback(prd_with_context)

call_kwargs = mock_agent.regenerate_with_feedback.call_args.kwargs
call_kwargs = mock_agent.regenerate_document_with_feedback.call_args.kwargs
assert call_kwargs["original_content"] == "# Original PRD"
assert call_kwargs["content_type"] == "prd"

Expand All @@ -243,15 +273,20 @@ async def test_feedback_is_passed_to_agent(self, prd_with_context):
mock_jira.update_description = AsyncMock()
mock_jira.add_comment = AsyncMock()
mock_jira.add_structured_comment = AsyncMock()
mock_jira.get_issue = AsyncMock(return_value=MagicMock(project_key="TEST"))
mock_jira.close = AsyncMock()

mock_agent = MagicMock()
mock_agent.regenerate_with_feedback = AsyncMock(return_value="# Revised")
mock_agent.regenerate_document_with_feedback = AsyncMock(
return_value=ArtifactDocument(content="# Revised", repositories=["acme/repo"])
)
mock_agent.close = AsyncMock()

with patch("forge.workflow.nodes.prd_generation.JiraClient", return_value=mock_jira):
with patch("forge.workflow.stations.artifact_generation.ForgeAgent", return_value=mock_agent):
with patch(
"forge.workflow.stations.artifact_generation.ForgeAgent", return_value=mock_agent
):
await regenerate_prd_with_feedback(prd_with_context)

call_kwargs = mock_agent.regenerate_with_feedback.call_args.kwargs
call_kwargs = mock_agent.regenerate_document_with_feedback.call_args.kwargs
assert "security" in call_kwargs["feedback"].lower()
Loading
Loading