Files
agentci/tests/workflow_support.py
StanPonomarev ce9f1e3d20
Publish container image / Build and push (push) Successful in 32s
refactor tests
2026-07-26 23:49:40 +02:00

179 lines
5.5 KiB
Python

from __future__ import annotations
from collections.abc import Iterable
from dataclasses import dataclass
from pathlib import Path
from types import SimpleNamespace
from typing import Any, cast
from pydantic import BaseModel
from agentci.engine.model import Workflow
from agentci.engine.run import JobRun
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 ScriptedOpenCode:
def __init__(
self,
responses: Iterable[BaseModel] = (),
trace: list[tuple[object, ...]] | None = None,
) -> None:
self.responses = list(responses)
self.created_sessions: list[tuple[Path, str]] = []
self.resume_calls: list[dict[str, Any]] = []
self.trace = trace if trace is not None else []
async def create_session(self, workspace: Path, title: str) -> str:
session_id = f"{title}-session"
self.created_sessions.append((workspace, title))
self.trace.append(("create_session", title, session_id))
return session_id
async def resume(self, **values: Any) -> BaseModel:
self.resume_calls.append(values)
self.trace.append(("resume", values["session_id"], values["result_type"]))
assert self.responses, "unexpected OpenCode resume call"
response = self.responses.pop(0)
assert isinstance(response, values["result_type"]), (
f"expected {values['result_type'].__name__}, got {type(response).__name__}"
)
return response
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 RecordingDevelopment:
description = "Python 3.13"
def __init__(self) -> None:
self.workspaces: list[Path] = []
async def prepare(self, workspace: Path) -> None:
self.workspaces.append(workspace)
class RecordingGit:
def __init__(
self,
*,
changed: bool = True,
clone_sha: str = "base-sha",
sync_sha: str = "head-sha",
commit_sha: str = "commit-sha",
) -> None:
self.changed = changed
self.clone_sha = clone_sha
self.sync_sha = sync_sha
self.commit_sha = commit_sha
self.calls: list[tuple[object, ...]] = []
async def clone(self, owner: str, repo: str, branch: str, destination: Path) -> str:
self.calls.append(("clone", owner, repo, branch, destination))
return self.clone_sha
async def create_branch(self, workspace: Path, branch: str) -> None:
self.calls.append(("create_branch", workspace, branch))
async def sync_branch(self, workspace: Path, branch: str) -> str:
self.calls.append(("sync_branch", workspace, branch))
return self.sync_sha
async def has_changes(self, workspace: Path) -> bool:
self.calls.append(("has_changes", workspace))
return self.changed
async def diff_check(self, workspace: Path) -> None:
self.calls.append(("diff_check", workspace))
async def commit(self, workspace: Path, message: str) -> str:
self.calls.append(("commit", workspace, message))
return self.commit_sha
async def push(self, workspace: Path, branch: str, *, set_upstream: bool = False) -> None:
self.calls.append(("push", workspace, branch, set_upstream))
@dataclass(frozen=True)
class WorkflowHarness:
services: WorkflowServices
run: RecordingRun
git: RecordingGit
development: RecordingDevelopment
prompts: RecordingPrompts
opencode: ScriptedOpenCode
trace: list[tuple[object, ...]]
def make_workflow_harness(
*,
settings: SimpleNamespace,
repository: object,
gitea: object,
responses: Iterable[BaseModel] = (),
trace: list[tuple[object, ...]] | None = None,
changed: bool = True,
clone_sha: str = "base-sha",
sync_sha: str = "head-sha",
commit_sha: str = "commit-sha",
) -> WorkflowHarness:
shared_trace = trace if trace is not None else []
run = RecordingRun()
git = RecordingGit(
changed=changed,
clone_sha=clone_sha,
sync_sha=sync_sha,
commit_sha=commit_sha,
)
development = RecordingDevelopment()
prompts = RecordingPrompts()
opencode = ScriptedOpenCode(responses, shared_trace)
services = cast(
WorkflowServices,
SimpleNamespace(
settings=settings,
repository=repository,
gitea=gitea,
git=git,
development=development,
prompts=prompts,
opencode=opencode,
),
)
return WorkflowHarness(
services=services,
run=run,
git=git,
development=development,
prompts=prompts,
opencode=opencode,
trace=shared_trace,
)