Files
agentci/tests/test_workflow_plan.py
T
2026-07-22 23:49:44 +02:00

440 lines
14 KiB
Python

import json
from pathlib import Path
from types import SimpleNamespace
from typing import Any, cast
import pytest
from pydantic import BaseModel
from agentci.engine.model import Job, JobKind, Workflow, WorkflowKind, WorkflowStatus
from agentci.engine.run import JobRun
from agentci.integrations.gitea.models import CommentInfo, IssueInfo, PullRequestInfo
from agentci.workflows.model import DiscussionReply, PlanArtifact, ReviewReport
from agentci.workflows.plan import create_plan, discuss_plan, iterate_plan
from agentci.workflows.render import JobRejected
from agentci.workflows.services import WorkflowServices
class RecordingRun(JobRun):
def __init__(self) -> None:
self.stages: list[str] = []
self.created_workflows: list[tuple[Workflow, str]] = []
self.linked_workflows: list[tuple[str, str]] = []
self.linked_sessions: list[str] = []
async def stage(self, stage: str) -> None:
self.stages.append(stage)
async def create_workflow(self, workflow: Workflow, stage: str) -> None:
self.created_workflows.append((workflow, stage))
async def link_workflow(self, workflow_id: str, stage: str) -> None:
self.linked_workflows.append((workflow_id, stage))
async def link_session(self, session_id: str) -> None:
self.linked_sessions.append(session_id)
class FakeRepository:
def __init__(
self,
*,
latest: Workflow | None = None,
implementations: list[Workflow] | None = None,
) -> None:
self.latest = latest
self.implementations = implementations or []
self.saved_workflows: list[Workflow] = []
self.latest_calls = 0
async def latest_workflow(self, *_args: object) -> Workflow | None:
self.latest_calls += 1
return self.latest
async def implementation_workflows(self, *_args: object) -> list[Workflow]:
return self.implementations
async def operational_comment_ids(self, *_args: object) -> set[int]:
return {91}
async def save_workflow(self, workflow: Workflow) -> None:
self.saved_workflows.append(workflow)
self.latest = workflow
class FakeGitea:
def __init__(self, pulls: dict[int, PullRequestInfo] | None = None) -> None:
self.pulls = pulls or {}
self.default_branch_calls: list[tuple[str, str]] = []
self.pull_calls: list[int] = []
async def default_branch(self, owner: str, repo: str) -> str:
self.default_branch_calls.append((owner, repo))
return "trunk"
async def issue(self, *_args: object) -> IssueInfo:
return IssueInfo(3, "Plan feature", "Build the feature.", "open")
async def issue_comments(self, *_args: object) -> list[CommentInfo]:
return [
CommentInfo(4, "alice", "Use the existing API.", "2026-07-01"),
CommentInfo(91, "agentci", "Job queued", "2026-07-02"),
]
async def pull_request(self, _owner: str, _repo: str, number: int) -> PullRequestInfo:
self.pull_calls.append(number)
return self.pulls[number]
class RecordingGit:
def __init__(self) -> None:
self.clone_calls: list[tuple[str, str, str, Path]] = []
async def clone(self, owner: str, repo: str, branch: str, destination: Path) -> str:
self.clone_calls.append((owner, repo, branch, destination))
return "base-sha"
class RecordingPrompts:
def __init__(self) -> None:
self.calls: list[tuple[str, dict[str, str]]] = []
def render(self, name: str, **values: str) -> str:
self.calls.append((name, values))
return f"rendered {name}"
class RecordingOpenCode:
def __init__(self, responses: list[BaseModel] | None = None) -> None:
self.responses = list(responses or [])
self.created_sessions: list[tuple[Path, str]] = []
self.resume_calls: list[dict[str, Any]] = []
async def create_session(self, workspace: Path, title: str) -> str:
self.created_sessions.append((workspace, title))
return f"{title}-session"
async def resume(self, **values: Any) -> BaseModel:
self.resume_calls.append(values)
response = self.responses.pop(0)
assert isinstance(response, values["result_type"])
return response
def job(kind: JobKind, *, message: str | None = "Please be specific.") -> Job:
return Job(
id="job-plan",
kind=kind,
target_key="org/repo:issue:3",
repo_owner="org",
repo_name="repo",
issue_number=3,
pr_number=None,
requester="alice",
comment_id=5,
delivery_id="delivery-plan",
receive_sequence=1,
command_body="/agent plan",
message=message,
)
def plan_workflow(
*,
runtime: str = "opencode",
primary_session_id: str | None = "primary-session",
reviewer_session_id: str | None = "reviewer-session",
artifact: str | None = "Original plan",
) -> Workflow:
return Workflow(
id="plan-flow",
kind=WorkflowKind.PLAN,
repo_owner="org",
repo_name="repo",
issue_number=3,
workspace_path=Path("/workspace/plan"),
base_sha="base-sha",
runtime=runtime,
primary_session_id=primary_session_id,
reviewer_session_id=reviewer_session_id,
artifact=artifact,
review_json='{"summary":"Prior","findings":[]}',
status=WorkflowStatus.COMPLETED,
)
def implementation_workflow(pr_number: int | None = 9) -> Workflow:
return Workflow(
id="implementation-flow",
kind=WorkflowKind.IMPLEMENT,
repo_owner="org",
repo_name="repo",
issue_number=3,
workspace_path=Path("/workspace/implementation"),
base_sha="base",
pr_number=pr_number,
status=WorkflowStatus.COMPLETED,
)
def pull(*, state: str = "open", merged: bool = False) -> PullRequestInfo:
return PullRequestInfo(
number=9,
title="Agent implementation",
body="Body",
state=state,
merged=merged,
base_branch="trunk",
head_branch="agent/feature",
head_sha="head",
head_owner="org",
head_repo="repo",
)
def make_services(
tmp_path: Path,
*,
latest: Workflow | None = None,
implementations: list[Workflow] | None = None,
pulls: dict[int, PullRequestInfo] | None = None,
responses: list[BaseModel] | None = None,
) -> tuple[
WorkflowServices,
FakeRepository,
FakeGitea,
RecordingGit,
RecordingPrompts,
RecordingOpenCode,
]:
repository = FakeRepository(latest=latest, implementations=implementations)
gitea = FakeGitea(pulls)
git = RecordingGit()
prompts = RecordingPrompts()
opencode = RecordingOpenCode(responses)
services = cast(
WorkflowServices,
SimpleNamespace(
settings=SimpleNamespace(
workspaces_dir=tmp_path / "workspaces",
plan_model="provider/model",
plan_variant="high",
plan_review_rounds=3,
),
repository=repository,
gitea=gitea,
git=git,
prompts=prompts,
opencode=opencode,
),
)
return services, repository, gitea, git, prompts, opencode
async def test_create_plan_completes_primary_and_independent_review(
tmp_path: Path,
) -> None:
services, repository, gitea, git, prompts, opencode = make_services(
tmp_path,
responses=[
PlanArtifact(plan_markdown="# Complete plan"),
ReviewReport(summary="Ready", findings=[]),
],
)
run = RecordingRun()
body = await create_plan(job(JobKind.PLAN), run, services)
created, stage = run.created_workflows[0]
assert stage == "planning"
assert created.kind is WorkflowKind.PLAN
assert created.workspace_path == tmp_path / "workspaces" / created.id / "repo"
assert git.clone_calls == [("org", "repo", "trunk", created.workspace_path)]
assert gitea.default_branch_calls == [("org", "repo")]
assert run.linked_sessions == ["plan-session"]
assert opencode.created_sessions == [
(created.workspace_path, "plan"),
(created.workspace_path, "plan-review"),
]
assert [call["result_type"] for call in opencode.resume_calls] == [
PlanArtifact,
ReviewReport,
]
assert prompts.calls[0][0] == "plan_initial"
assert prompts.calls[0][1]["request"] == "Please be specific."
assert "Use the existing API." in prompts.calls[0][1]["context"]
assert "Job queued" not in prompts.calls[0][1]["context"]
completed = repository.saved_workflows[-1]
assert completed.status is WorkflowStatus.COMPLETED
assert completed.primary_session_id == "plan-session"
assert completed.reviewer_session_id == "plan-review-session"
assert completed.artifact == "# Complete plan"
assert json.loads(completed.review_json or "") == {
"summary": "Ready",
"findings": [],
}
assert body == f"<!-- agentci:plan workflow={created.id} -->\n# Complete plan"
async def test_discuss_plan_resumes_primary_session_without_replacing_artifact(
tmp_path: Path,
) -> None:
existing = plan_workflow()
services, repository, _, _, prompts, opencode = make_services(
tmp_path,
latest=existing,
responses=[DiscussionReply(markdown="The API remains compatible.")],
)
run = RecordingRun()
body = await discuss_plan(job(JobKind.DISCUSS), run, services)
assert run.linked_workflows == [(existing.id, "discussing")]
assert opencode.created_sessions == []
assert opencode.resume_calls[0]["session_id"] == "primary-session"
assert opencode.resume_calls[0]["schema_name"] == "discussion.json"
assert prompts.calls == [
(
"discuss",
{"artifact": "Original plan", "message": "Please be specific."},
)
]
assert repository.saved_workflows == []
assert body == ("<!-- agentci:discussion workflow=plan-flow -->\nThe API remains compatible.")
@pytest.mark.parametrize(
("latest", "message"),
[
(None, "No completed plan exists. Start with `/agent plan`."),
(
plan_workflow(runtime="codex"),
"The latest plan predates OpenCode and cannot be resumed; start a new `/agent plan`.",
),
(
plan_workflow(primary_session_id=None),
"The latest plan cannot be resumed; start a new `/agent plan`.",
),
(
plan_workflow(artifact=None),
"The latest plan cannot be resumed; start a new `/agent plan`.",
),
],
)
async def test_discuss_plan_rejects_missing_or_incompatible_plan(
tmp_path: Path,
latest: Workflow | None,
message: str,
) -> None:
services, _, _, _, _, opencode = make_services(tmp_path, latest=latest)
with pytest.raises(JobRejected) as error:
await discuss_plan(job(JobKind.DISCUSS), RecordingRun(), services)
assert str(error.value) == message
assert opencode.resume_calls == []
async def test_iterate_plan_revises_then_reuses_reviewer_session(
tmp_path: Path,
) -> None:
existing = plan_workflow()
services, repository, gitea, _, prompts, opencode = make_services(
tmp_path,
latest=existing,
implementations=[implementation_workflow(), implementation_workflow(None)],
pulls={9: pull(state="closed")},
responses=[
PlanArtifact(plan_markdown="# Revised plan"),
ReviewReport(summary="Ready", findings=[]),
],
)
run = RecordingRun()
body = await iterate_plan(job(JobKind.ITERATE_PLAN, message=None), run, services)
assert gitea.pull_calls == [9]
assert run.linked_workflows == [(existing.id, "iterating plan")]
assert opencode.created_sessions == []
assert [call["session_id"] for call in opencode.resume_calls] == [
"primary-session",
"reviewer-session",
]
assert [call["result_type"] for call in opencode.resume_calls] == [
PlanArtifact,
ReviewReport,
]
iterate_prompt = prompts.calls[0]
assert iterate_prompt[0] == "plan_iterate"
assert iterate_prompt[1]["artifact"] == "Original plan"
assert iterate_prompt[1]["message"] == ("(refine using the latest discussion and prior review)")
assert json.loads(iterate_prompt[1]["review"]) == {
"summary": "Prior",
"findings": [],
}
completed = repository.saved_workflows[-1]
assert completed.artifact == "# Revised plan"
assert completed.status is WorkflowStatus.COMPLETED
assert body == "<!-- agentci:plan workflow=plan-flow -->\n# Revised plan"
@pytest.mark.parametrize(
("latest", "message"),
[
(None, "No completed plan exists. Start with `/agent plan`."),
(
plan_workflow(runtime="codex"),
"The latest plan predates OpenCode; start a new plan.",
),
(
plan_workflow(primary_session_id=None),
"The latest plan is missing resumable sessions; start a new plan.",
),
(
plan_workflow(reviewer_session_id=None),
"The latest plan is missing resumable sessions; start a new plan.",
),
(plan_workflow(artifact=None), "The latest plan has no saved artifact."),
],
)
async def test_iterate_plan_rejects_missing_or_incompatible_plan(
tmp_path: Path,
latest: Workflow | None,
message: str,
) -> None:
services, repository, _, _, _, opencode = make_services(tmp_path, latest=latest)
with pytest.raises(JobRejected) as error:
await iterate_plan(job(JobKind.ITERATE_PLAN), RecordingRun(), services)
assert str(error.value) == message
assert repository.saved_workflows == []
assert opencode.resume_calls == []
@pytest.mark.parametrize(
"blocking_pull",
[pull(state="open"), pull(state="closed", merged=True)],
ids=["open", "merged"],
)
async def test_iterate_plan_rejects_after_agent_pr_is_open_or_merged(
tmp_path: Path, blocking_pull: PullRequestInfo
) -> None:
services, repository, _, _, _, opencode = make_services(
tmp_path,
latest=plan_workflow(),
implementations=[implementation_workflow()],
pulls={9: blocking_pull},
)
with pytest.raises(JobRejected) as error:
await iterate_plan(job(JobKind.ITERATE_PLAN), RecordingRun(), services)
assert str(error.value) == (
"Issue plan iteration is disabled because agent PR #9 is open or merged. "
"Iterate an open implementation on its PR."
)
assert repository.latest_calls == 0
assert opencode.resume_calls == []