diff --git a/roboco/services/gateway/choreographer.py b/roboco/services/gateway/choreographer.py index 8a0eedbd..592a0fab 100644 --- a/roboco/services/gateway/choreographer.py +++ b/roboco/services/gateway/choreographer.py @@ -255,14 +255,14 @@ class Choreographer: # Resumption after QA rejection: agent already owned the task, # role-typed claim already passed at original claim time. if t.assigned_to != agent_id: - t = await self.task.claim(agent_id, task_id) - t = await self.task.start(agent_id, task_id) + t = await self.task.claim(task_id, agent_id) + t = await self.task.start(task_id, agent_id) elif status == "pending": # Fresh claim — run all claim-time gates BEFORE mutating state. if guard := await self._run_claim_guards(agent_id=agent_id, task=t): return self._with_briefing(guard, briefing) if t.assigned_to is None or t.assigned_to != agent_id: - t = await self.task.claim(agent_id, task_id) + t = await self.task.claim(task_id, agent_id) if not t.plan and not plan: remediate = ( f"call i_will_work_on(task_id='{task_id}'," @@ -275,7 +275,7 @@ class Choreographer: ) if plan: t = await self.task.set_plan(task_id, plan) - t = await self.task.start(agent_id, task_id) + t = await self.task.start(task_id, agent_id) elif status == "claimed" and t.assigned_to == agent_id: # Resumption: skip sibling-sequence (already passed at claim). # Still enforce already_active/paused so concurrent claims fail. @@ -284,7 +284,7 @@ class Choreographer: ) if guard: return self._with_briefing(guard, briefing) - t = await self.task.start(agent_id, task_id) + t = await self.task.start(task_id, agent_id) else: return Envelope.invalid_state( message=f"task {task_id} is in {status}; cannot start work", @@ -1045,7 +1045,7 @@ class Choreographer: return rejection if t.assigned_to is None or t.assigned_to != pm_agent_id: - t = await self.task.claim(pm_agent_id, task_id) + t = await self.task.claim(task_id, pm_agent_id) if t is None: return Envelope.invalid_state( message="claim failed", @@ -1053,7 +1053,7 @@ class Choreographer: context_briefing=await self._briefing_for(pm_agent_id, task_id), ) await self.task.set_plan(task_id, plan) - t = await self.task.start(pm_agent_id, task_id) + t = await self.task.start(task_id, pm_agent_id) return Envelope.ok( status=str(t.status) if t else "in_progress", task_id=str(task_id), diff --git a/tests/unit/gateway/test_choreographer_claim_guards.py b/tests/unit/gateway/test_choreographer_claim_guards.py index 500ee0cd..52e5dc65 100644 --- a/tests/unit/gateway/test_choreographer_claim_guards.py +++ b/tests/unit/gateway/test_choreographer_claim_guards.py @@ -147,7 +147,7 @@ async def test_i_will_work_on_allows_when_earlier_sibling_terminal() -> None: env = await c.i_will_work_on(agent_id, target_id) assert env.error is None - task_svc.claim.assert_awaited_once_with(agent_id, target_id) + task_svc.claim.assert_awaited_once_with(target_id, agent_id) @pytest.mark.asyncio @@ -242,7 +242,7 @@ async def test_i_will_work_on_resumption_does_not_self_block() -> None: env = await c.i_will_work_on(agent_id, task_id) assert env.error is None - task_svc.start.assert_awaited_once_with(agent_id, task_id) + task_svc.start.assert_awaited_once_with(task_id, agent_id) # --------------------------------------------------------------------------- diff --git a/tests/unit/gateway/test_choreographer_dev.py b/tests/unit/gateway/test_choreographer_dev.py index 0e8b52b6..cdc5ca97 100644 --- a/tests/unit/gateway/test_choreographer_dev.py +++ b/tests/unit/gateway/test_choreographer_dev.py @@ -120,9 +120,9 @@ async def test_i_will_work_on_pending_with_plan() -> None: env = await c.i_will_work_on(agent_id, task_id, plan="do x then y") assert env.error is None assert env.status == "in_progress" - task_svc.claim.assert_awaited_once_with(agent_id, task_id) + task_svc.claim.assert_awaited_once_with(task_id, agent_id) task_svc.set_plan.assert_awaited_once() - task_svc.start.assert_awaited_once_with(agent_id, task_id) + task_svc.start.assert_awaited_once_with(task_id, agent_id) @pytest.mark.asyncio @@ -177,7 +177,7 @@ async def test_i_will_work_on_needs_revision_re_starts() -> None: env = await c.i_will_work_on(agent_id, task_id) assert env.status == "in_progress" task_svc.claim.assert_not_awaited() # already assigned - task_svc.start.assert_awaited_once_with(agent_id, task_id) + task_svc.start.assert_awaited_once_with(task_id, agent_id) @pytest.mark.asyncio diff --git a/tests/unit/gateway/test_choreographer_pm_extras.py b/tests/unit/gateway/test_choreographer_pm_extras.py index 5c825978..7d166614 100644 --- a/tests/unit/gateway/test_choreographer_pm_extras.py +++ b/tests/unit/gateway/test_choreographer_pm_extras.py @@ -90,9 +90,9 @@ async def test_i_will_plan_claims_starts_and_sets_plan() -> None: env = await c.i_will_plan(pm_id, task_id, plan="break the work into 3 subtasks") assert env.error is None assert env.status == "in_progress" - task_svc.claim.assert_awaited_once_with(pm_id, task_id) + task_svc.claim.assert_awaited_once_with(task_id, pm_id) task_svc.set_plan.assert_awaited_once() - task_svc.start.assert_awaited_once_with(pm_id, task_id) + task_svc.start.assert_awaited_once_with(task_id, pm_id) @pytest.mark.asyncio diff --git a/tests/unit/gateway/test_claim_arg_order.py b/tests/unit/gateway/test_claim_arg_order.py new file mode 100644 index 00000000..8dfda23a --- /dev/null +++ b/tests/unit/gateway/test_claim_arg_order.py @@ -0,0 +1,217 @@ +"""Pin: Choreographer must call task.claim/start with (task_id, agent_id). + +Service signatures in ``roboco/services/task.py`` are: + + async def claim(self, task_id: UUID, agent_id: UUID, ...) -> TaskTable | None + async def start(self, task_id: UUID, agent_id: UUID | None = None, ...) -> ... + +Earlier choreographer code passed (agent_id, task_id) — when the SQL lookup +ran ``WHERE id = ``, no row matched and the call returned None, +which the choreographer interpreted as "task unchanged". The whole gateway +claim path was non-functional against a real DB. Existing unit tests pinned +the buggy order so the bug stayed invisible. + +This test pins the correct order at every Choreographer call site that +forwards into the service. +""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import AsyncMock, MagicMock +from uuid import uuid4 + +import pytest +from roboco.services.gateway.choreographer import Choreographer, ChoreographerDeps + + +def _make_deps(**overrides: Any) -> ChoreographerDeps: + base = { + "task": AsyncMock(), + "work_session": AsyncMock(), + "git": AsyncMock(), + "a2a": AsyncMock(), + "journal": AsyncMock(), + "audit": AsyncMock(), + "evidence_repo": AsyncMock(), + } + base.update(overrides) + repo = base["evidence_repo"] + for method in ( + "list_unread_a2a", + "list_unread_mentions", + "list_pending_notifications", + "task_metadata_gaps", + "recent_team_activity", + "blockers_in_lane", + "journal_highlights_for_task", + ): + getattr(repo, method).return_value = [] + return ChoreographerDeps(**base) + + +@pytest.mark.asyncio +async def test_i_will_work_on_pending_calls_claim_with_task_id_first() -> None: + """Dev pending claim: positional args must be (task_id, agent_id).""" + agent_id = uuid4() + task_id = uuid4() + pending = MagicMock( + id=task_id, + status="pending", + plan=None, + assigned_to=None, + task_type="code", + parent_task_id=None, + sequence=0, + team="backend", + ) + claimed = MagicMock( + id=task_id, + status="claimed", + plan=None, + assigned_to=agent_id, + task_type="code", + ) + with_plan = MagicMock( + id=task_id, + status="claimed", + plan={"text": "x"}, + assigned_to=agent_id, + task_type="code", + ) + started = MagicMock( + id=task_id, + status="in_progress", + plan={"text": "x"}, + assigned_to=agent_id, + task_type="code", + ) + task_svc = AsyncMock() + task_svc.get.return_value = pending + task_svc.agent_for.return_value = MagicMock(role="developer", team="backend") + task_svc.list_in_progress_for_agent.return_value = [] + task_svc.list_paused_for_agent.return_value = [] + task_svc.get_subtasks.return_value = [] + task_svc.claim.return_value = claimed + task_svc.set_plan.return_value = with_plan + task_svc.start.return_value = started + deps = _make_deps(task=task_svc) + c = Choreographer(deps) + + env = await c.i_will_work_on(agent_id, task_id, plan="x") + + # Service signature is (task_id, agent_id, ...) — pin that order. + task_svc.claim.assert_awaited_once_with(task_id, agent_id) + task_svc.start.assert_awaited_once_with(task_id, agent_id) + assert env.error is None + + +@pytest.mark.asyncio +async def test_i_will_work_on_needs_revision_calls_start_with_task_id_first() -> None: + """Dev needs_revision resume: start args must be (task_id, agent_id).""" + agent_id = uuid4() + task_id = uuid4() + nr = MagicMock( + id=task_id, + status="needs_revision", + assigned_to=agent_id, + plan={"x": 1}, + task_type="code", + ) + started = MagicMock( + id=task_id, status="in_progress", assigned_to=agent_id, plan={"x": 1} + ) + task_svc = AsyncMock() + task_svc.get.return_value = nr + task_svc.start.return_value = started + deps = _make_deps(task=task_svc) + c = Choreographer(deps) + + env = await c.i_will_work_on(agent_id, task_id) + + task_svc.start.assert_awaited_once_with(task_id, agent_id) + task_svc.claim.assert_not_awaited() + assert env.error is None + + +@pytest.mark.asyncio +async def test_i_will_work_on_claimed_resumption_calls_start_with_task_id_first() -> ( + None +): + """Dev claimed resumption: start args must be (task_id, agent_id).""" + agent_id = uuid4() + task_id = uuid4() + claimed = MagicMock( + id=task_id, + status="claimed", + plan={"x": 1}, + assigned_to=agent_id, + parent_task_id=None, + sequence=0, + task_type="code", + team="backend", + branch_name="feature/backend/abc", + ) + started = MagicMock( + id=task_id, status="in_progress", plan={"x": 1}, assigned_to=agent_id + ) + task_svc = AsyncMock() + task_svc.get.return_value = claimed + task_svc.agent_for.return_value = MagicMock(role="developer", team="backend") + task_svc.list_in_progress_for_agent.return_value = [] + task_svc.list_paused_for_agent.return_value = [] + task_svc.get_subtasks.return_value = [] + task_svc.start.return_value = started + deps = _make_deps(task=task_svc) + c = Choreographer(deps) + + env = await c.i_will_work_on(agent_id, task_id) + + task_svc.start.assert_awaited_once_with(task_id, agent_id) + assert env.error is None + + +@pytest.mark.asyncio +async def test_i_will_plan_calls_claim_and_start_with_task_id_first() -> None: + """PM plan path: both claim and start must use (task_id, agent_id).""" + pm_id = uuid4() + task_id = uuid4() + pending = MagicMock( + id=task_id, + status="pending", + plan=None, + assigned_to=None, + task_type="planning", + parent_task_id=None, + sequence=0, + ) + claimed = MagicMock( + id=task_id, + status="claimed", + plan=None, + assigned_to=pm_id, + task_type="planning", + ) + started = MagicMock( + id=task_id, + status="in_progress", + plan={"text": "x"}, + assigned_to=pm_id, + task_type="planning", + ) + task_svc = AsyncMock() + task_svc.get.return_value = pending + task_svc.agent_for.return_value = MagicMock(role="cell_pm", team="backend") + task_svc.list_in_progress_for_agent.return_value = [] + task_svc.list_paused_for_agent.return_value = [] + task_svc.claim.return_value = claimed + task_svc.set_plan.return_value = claimed + task_svc.start.return_value = started + deps = _make_deps(task=task_svc) + c = Choreographer(deps) + + env = await c.i_will_plan(pm_id, task_id, plan="break the work into 3 subtasks") + + task_svc.claim.assert_awaited_once_with(task_id, pm_id) + task_svc.start.assert_awaited_once_with(task_id, pm_id) + assert env.error is None