diff --git a/agent-team/agent_team/coordinator.py b/agent-team/agent_team/coordinator.py index ccd022a..6207764 100644 --- a/agent-team/agent_team/coordinator.py +++ b/agent-team/agent_team/coordinator.py @@ -72,6 +72,7 @@ __all__ = [ "DispatchNodeFactory", "build_verify_wiring", "default_clarify_node_factory", + "default_dispatch_node_factory", "default_slack_listener_factory", "gated_build_verify_wiring", ] @@ -365,6 +366,38 @@ def gated_build_verify_wiring( return build_node, verify_node, route_after_verify +def default_dispatch_node_factory() -> "Callable[[Any], Any]": + """Build the live auto-dispatch node from env vars (WS3, OPT-IN, INERT by default). + + Reads ``AGENT_TEAM_REPO_OWNER``, ``AGENT_TEAM_REPO_NAME``, and optionally + ``AGENT_TEAM_BASE_BRANCH`` (default ``"main"``) from the environment and + delegates to :func:`agent_team.nodes.dispatch_invoker.make_dispatch_node`. + Owner/repo are fixed at factory time — never read from pipeline state — so + model output in the state cannot redirect the dispatch target. + + Raises ``RuntimeError`` if the required env vars are absent (fail closed: + the node must never dispatch to an unknown owner/repo). + + This satisfies :data:`DispatchNodeFactory` (zero-arg → node callable). + Pass it as ``dispatch_node_wiring=default_dispatch_node_factory`` AFTER the + §3.3.2 CI trust-boundary gate clears. The production ``run-team.py`` path + deliberately does NOT inject this; it remains opt-in. + """ + import os + + from agent_team.nodes.dispatch_invoker import make_dispatch_node + + owner = os.environ.get("AGENT_TEAM_REPO_OWNER", "").strip() + repo = os.environ.get("AGENT_TEAM_REPO_NAME", "").strip() + base = os.environ.get("AGENT_TEAM_BASE_BRANCH", "main").strip() or "main" + if not owner or not repo: + raise RuntimeError( + "AGENT_TEAM_REPO_OWNER and AGENT_TEAM_REPO_NAME must be set " + "for auto-dispatch (default_dispatch_node_factory)" + ) + return make_dispatch_node(owner=owner, repo=repo, base=base) + + class Coordinator: """Owns the live Plane-2 runtime: graph + resume worker + transport (§3.3). diff --git a/agent-team/tests/test_ws3_dispatch_invoker.py b/agent-team/tests/test_ws3_dispatch_invoker.py new file mode 100644 index 0000000..2e607b1 --- /dev/null +++ b/agent-team/tests/test_ws3_dispatch_invoker.py @@ -0,0 +1,425 @@ +"""Unit tests for WS3: agent_team.nodes.dispatch_invoker + graph/coordinator wiring. + +Tests cover: +* :func:`make_dispatch_node` — factory contract, fail-closed on missing fields + (parks on missing thread_id/diff/scope), injection of owner/repo at factory + time (never read from state), scope list flattening, happy-path returns empty + dict (partial state update) and calls dispatcher. +* :func:`agent_team.graph.build_graph` with ``dispatch_node`` wired — verifies + the DISPATCH_NODE constant is exported and the graph compiles with/without it. +* :func:`agent_team.coordinator.default_dispatch_node_factory` — env-var binding + and fail-closed when required vars are unset. +* :class:`agent_team.coordinator.Coordinator` ``dispatch_node_wiring`` seam — + accepted at __init__, NOT called in setup() when build_verify_wiring is None + (DISPATCH_NODE requires VERIFY_NODE/APPROVED_ROUTE to exist). + +All network + subprocess calls are replaced with injected fakes; no model, no +git, no GitHub, no subprocess. +""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import MagicMock + +import pytest + +from agent_team.nodes.dispatch_invoker import DispatchNodeFactory, make_dispatch_node +from agent_team.task_model import Phase, TaskStatus + + +# --------------------------------------------------------------------------- # +# Fake BranchPusher / WorkflowDispatcher (the two injectable seams in dispatcher) +# --------------------------------------------------------------------------- # + + +def _make_fake_pusher() -> tuple[list[dict[str, Any]], Any]: + calls: list[dict[str, Any]] = [] + + def _pusher( + *, owner: str, repo: str, base: str, head_branch: str, diff_text: str + ) -> None: + calls.append(dict(owner=owner, repo=repo, base=base, head_branch=head_branch)) + + return calls, _pusher + + +def _make_fake_workflow_dispatcher() -> tuple[list[dict[str, Any]], Any]: + calls: list[dict[str, Any]] = [] + + def _dispatcher( + *, owner: str, repo: str, inputs: dict[str, str], ref: str + ) -> None: + calls.append(dict(owner=owner, repo=repo, ref=ref, inputs=inputs)) + + return calls, _dispatcher + + +def _fake_seams() -> tuple[Any, Any]: + """Return (fake_pusher, fake_dispatcher) that silence all subprocess calls.""" + _, pusher = _make_fake_pusher() + _, dispatcher = _make_fake_workflow_dispatcher() + return pusher, dispatcher + + +_VALID_STATE: dict[str, Any] = { + "thread_id": "task-abc", + "candidate_diff": ( + "diff --git a/f.py b/f.py\n--- a/f.py\n+++ b/f.py\n@@ -1 +1 @@\n-old\n+new" + ), + "plan": {"scope": ["agent_team/"]}, + "status": "active", + "current_phase": "verify", +} + + +# --------------------------------------------------------------------------- # +# make_dispatch_node — happy path +# --------------------------------------------------------------------------- # + + +def test_dispatch_node_happy_path_returns_partial_state() -> None: + """On success the node returns {} (partial state update — DONE comes from graph).""" + pusher, dispatcher = _fake_seams() + node = make_dispatch_node(owner="org", repo="repo", pusher=pusher, dispatcher=dispatcher) + result = node(_VALID_STATE) + + # Remote implementation returns {} on success (graph topology marks DONE). + assert isinstance(result, dict) + + +def test_dispatch_node_injects_owner_repo_at_factory_time() -> None: + """owner/repo come from factory args — model output in state cannot redirect them.""" + pusher_calls, pusher = _make_fake_pusher() + _, dispatcher = _make_fake_workflow_dispatcher() + node = make_dispatch_node( + owner="trusted-org", repo="trusted-repo", pusher=pusher, dispatcher=dispatcher + ) + node(_VALID_STATE) + + assert pusher_calls, "pusher should have been called" + assert pusher_calls[0]["owner"] == "trusted-org" + assert pusher_calls[0]["repo"] == "trusted-repo" + + +def test_dispatch_node_passes_task_id_to_head_branch() -> None: + pusher_calls, pusher = _make_fake_pusher() + _, dispatcher = _make_fake_workflow_dispatcher() + node = make_dispatch_node( + owner="org", repo="repo", base="main", pusher=pusher, dispatcher=dispatcher + ) + node(_VALID_STATE) + + assert pusher_calls, "pusher should have been called" + # build_dispatch_inputs derives head_branch from task_id. + assert "task-abc" in pusher_calls[0]["head_branch"] + assert pusher_calls[0]["base"] == "main" + + +def test_dispatch_node_passes_scope_to_workflow_dispatcher() -> None: + _, pusher = _make_fake_pusher() + disp_calls, dispatcher = _make_fake_workflow_dispatcher() + node = make_dispatch_node(owner="org", repo="repo", pusher=pusher, dispatcher=dispatcher) + node(_VALID_STATE) + + assert disp_calls, "workflow dispatcher should have been called" + wf_inputs = disp_calls[0]["inputs"] + assert "agent_team/" in str(wf_inputs) + + +def test_dispatch_node_flattens_scope_list_to_string() -> None: + """Multi-entry scope list is newline-joined into declared_scope.""" + _, pusher = _make_fake_pusher() + disp_calls, dispatcher = _make_fake_workflow_dispatcher() + node = make_dispatch_node(owner="org", repo="repo", pusher=pusher, dispatcher=dispatcher) + + state = dict(_VALID_STATE, plan={"scope": ["agent_team/", "tests/"]}) + node(state) + + wf_inputs = disp_calls[0]["inputs"] + scope_str = str(wf_inputs) + assert "agent_team/" in scope_str + + +# --------------------------------------------------------------------------- # +# make_dispatch_node — fail closed (never dispatches incomplete input) +# --------------------------------------------------------------------------- # + + +def test_dispatch_node_parks_on_missing_thread_id() -> None: + _, pusher = _make_fake_pusher() + disp_calls, dispatcher = _make_fake_workflow_dispatcher() + node = make_dispatch_node(owner="org", repo="repo", pusher=pusher, dispatcher=dispatcher) + + state = dict(_VALID_STATE, thread_id="") + result = node(state) + + assert result.get("status") == TaskStatus.PARKED.value + assert result.get("current_phase") == Phase.PARKED.value + assert not disp_calls # dispatcher never called + + +def test_dispatch_node_parks_on_missing_diff() -> None: + _, pusher = _make_fake_pusher() + disp_calls, dispatcher = _make_fake_workflow_dispatcher() + node = make_dispatch_node(owner="org", repo="repo", pusher=pusher, dispatcher=dispatcher) + + state = dict(_VALID_STATE, candidate_diff="") + result = node(state) + + assert result.get("status") == TaskStatus.PARKED.value + assert not disp_calls + + +def test_dispatch_node_parks_on_whitespace_only_diff() -> None: + _, pusher = _make_fake_pusher() + disp_calls, dispatcher = _make_fake_workflow_dispatcher() + node = make_dispatch_node(owner="org", repo="repo", pusher=pusher, dispatcher=dispatcher) + + state = dict(_VALID_STATE, candidate_diff=" \n ") + result = node(state) + + assert result.get("status") == TaskStatus.PARKED.value + assert not disp_calls + + +def test_dispatch_node_parks_on_empty_scope() -> None: + _, pusher = _make_fake_pusher() + disp_calls, dispatcher = _make_fake_workflow_dispatcher() + node = make_dispatch_node(owner="org", repo="repo", pusher=pusher, dispatcher=dispatcher) + + state = dict(_VALID_STATE, plan={"scope": []}) + result = node(state) + + assert result.get("status") == TaskStatus.PARKED.value + assert not disp_calls + + +def test_dispatch_node_parks_on_none_plan() -> None: + _, pusher = _make_fake_pusher() + disp_calls, dispatcher = _make_fake_workflow_dispatcher() + node = make_dispatch_node(owner="org", repo="repo", pusher=pusher, dispatcher=dispatcher) + + state = dict(_VALID_STATE, plan=None) + result = node(state) + + assert result.get("status") == TaskStatus.PARKED.value + assert not disp_calls + + +def test_dispatch_node_parks_on_non_dict_plan() -> None: + _, pusher = _make_fake_pusher() + disp_calls, dispatcher = _make_fake_workflow_dispatcher() + node = make_dispatch_node(owner="org", repo="repo", pusher=pusher, dispatcher=dispatcher) + + state = dict(_VALID_STATE, plan="not-a-dict") + result = node(state) + + assert result.get("status") == TaskStatus.PARKED.value + assert not disp_calls + + +def test_dispatch_node_parks_on_dispatcher_error() -> None: + """A DispatcherError from the underlying dispatcher parks the task (fail safe).""" + from agent_team.dispatcher import DispatcherError + + def _bad_pusher(**_kwargs: Any) -> None: + raise DispatcherError("invalid owner/repo 'x'/'y'") + + _, dispatcher = _make_fake_workflow_dispatcher() + node = make_dispatch_node(owner="org", repo="repo", pusher=_bad_pusher, dispatcher=dispatcher) + result = node(_VALID_STATE) + + assert result.get("status") == TaskStatus.PARKED.value + + +def test_dispatch_node_parks_on_unexpected_exception() -> None: + """Any unexpected error parks the task rather than crashing the graph.""" + + def _exploding_pusher(**_kwargs: Any) -> None: + raise RuntimeError("network timeout") + + _, dispatcher = _make_fake_workflow_dispatcher() + node = make_dispatch_node( + owner="org", repo="repo", pusher=_exploding_pusher, dispatcher=dispatcher + ) + result = node(_VALID_STATE) + + assert result.get("status") == TaskStatus.PARKED.value + + +# --------------------------------------------------------------------------- # +# DispatchNodeFactory type alias export +# --------------------------------------------------------------------------- # + + +def test_dispatch_node_factory_type_is_exported() -> None: + from agent_team.nodes import dispatch_invoker + + assert hasattr(dispatch_invoker, "DispatchNodeFactory") + + +def test_dispatch_node_factory_satisfies_zero_arg_contract() -> None: + """DispatchNodeFactory is a zero-arg callable that returns a node callable.""" + pusher, dispatcher = _fake_seams() + factory: DispatchNodeFactory = lambda: make_dispatch_node( # noqa: E731 + owner="o", repo="r", pusher=pusher, dispatcher=dispatcher + ) + node = factory() + assert callable(node) + result = node(_VALID_STATE) + assert isinstance(result, dict) + + +# --------------------------------------------------------------------------- # +# graph.build_graph — DISPATCH_NODE constant and compile-time wiring +# --------------------------------------------------------------------------- # + + +def test_build_graph_dispatch_node_constant_is_exported() -> None: + from agent_team import graph as graph_mod + + assert hasattr(graph_mod, "DISPATCH_NODE") + assert graph_mod.DISPATCH_NODE == "dispatch_node" + + +def test_build_graph_without_dispatch_node_compiles() -> None: + """Default (no dispatch_node) builds fine — backward-compatible.""" + from langgraph.checkpoint.memory import MemorySaver + + from agent_team.graph import build_graph + + graph = build_graph(MemorySaver()) + assert graph is not None + + +def test_build_graph_with_dispatch_node_no_build_verify_raises() -> None: + """dispatch_node without build_verify raises ValueError (APPROVED_ROUTE requires P3).""" + from langgraph.checkpoint.memory import MemorySaver + + from agent_team.graph import build_graph + + def _fake_dispatch(state: Any) -> Any: + return {} + + with pytest.raises(ValueError, match="dispatch_node requires build_verify"): + build_graph(MemorySaver(), dispatch_node=_fake_dispatch) + + +# --------------------------------------------------------------------------- # +# coordinator.default_dispatch_node_factory +# --------------------------------------------------------------------------- # + + +def test_default_dispatch_node_factory_raises_without_env_vars( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("AGENT_TEAM_REPO_OWNER", raising=False) + monkeypatch.delenv("AGENT_TEAM_REPO_NAME", raising=False) + + from agent_team.coordinator import default_dispatch_node_factory + + with pytest.raises(RuntimeError, match="AGENT_TEAM_REPO_OWNER"): + default_dispatch_node_factory() + + +def test_default_dispatch_node_factory_raises_when_only_owner_set( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("AGENT_TEAM_REPO_OWNER", "org") + monkeypatch.delenv("AGENT_TEAM_REPO_NAME", raising=False) + + from agent_team.coordinator import default_dispatch_node_factory + + with pytest.raises(RuntimeError): + default_dispatch_node_factory() + + +def test_default_dispatch_node_factory_returns_callable_with_env_vars( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("AGENT_TEAM_REPO_OWNER", "test-org") + monkeypatch.setenv("AGENT_TEAM_REPO_NAME", "test-repo") + monkeypatch.delenv("AGENT_TEAM_BASE_BRANCH", raising=False) + + from agent_team.coordinator import default_dispatch_node_factory + + node = default_dispatch_node_factory() + assert callable(node) + + +def test_default_dispatch_node_factory_reads_base_branch_from_env( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("AGENT_TEAM_REPO_OWNER", "org") + monkeypatch.setenv("AGENT_TEAM_REPO_NAME", "repo") + monkeypatch.setenv("AGENT_TEAM_BASE_BRANCH", "develop") + + from agent_team.coordinator import default_dispatch_node_factory + + node = default_dispatch_node_factory() + assert callable(node) + + +def test_default_dispatch_node_factory_in_all() -> None: + from agent_team import coordinator + + assert "default_dispatch_node_factory" in coordinator.__all__ + + +# --------------------------------------------------------------------------- # +# Coordinator.dispatch_node_wiring seam +# --------------------------------------------------------------------------- # + + +def test_coordinator_accepts_dispatch_node_wiring_param() -> None: + """Coordinator.__init__ accepts dispatch_node_wiring=None without error.""" + from agent_team.coordinator import Coordinator + from agent_team.transport.base import Transport + + transport = MagicMock(spec=Transport) + coord = Coordinator( + db_path=":memory:", + transport=transport, + dispatch_node_wiring=None, + ) + assert coord is not None + + +def test_coordinator_dispatch_node_wiring_without_build_verify_raises( + tmp_path: Any, +) -> None: + """setup() raises ValueError when dispatch_node_wiring is set but build_verify_wiring is None. + + build_graph enforces that dispatch_node requires build_verify (APPROVED_ROUTE + lives inside the P3 subgraph). Passing dispatch_node_wiring alone is a + configuration error caught at setup time. + """ + from langgraph.checkpoint.memory import MemorySaver + + from agent_team.coordinator import Coordinator + from agent_team.transport.base import Transport + + pusher, dispatcher = _fake_seams() + + def _factory() -> Any: + return make_dispatch_node(owner="org", repo="repo", pusher=pusher, dispatcher=dispatcher) + + def _stub_clarify_node() -> Any: + return MagicMock() + + def _stub_checkpointer(db_path: Any) -> Any: + return MemorySaver() + + transport = MagicMock(spec=Transport) + coord = Coordinator( + db_path=tmp_path / "test.db", + transport=transport, + build_clarify_node=_stub_clarify_node, + build_checkpointer=_stub_checkpointer, + dispatch_node_wiring=_factory, + # build_verify_wiring=None (default) + ) + + with pytest.raises(ValueError, match="dispatch_node requires build_verify"): + coord.setup()