feat(cloud): skip MCP-busy devices in scheduler

This commit is contained in:
2026-07-21 14:37:21 +08:00
parent 9e3007e7f6
commit b0932dd398
8 changed files with 172 additions and 1 deletions
@@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
from sqlalchemy import ( from sqlalchemy import (
Boolean,
ForeignKey, ForeignKey,
Index, Index,
Integer, Integer,
@@ -86,6 +87,7 @@ class PooledDeviceRow(Base):
status: Mapped[str] = mapped_column(String, nullable=False) status: Mapped[str] = mapped_column(String, nullable=False)
capability_tags_json: Mapped[str] = mapped_column(Text, nullable=False) capability_tags_json: Mapped[str] = mapped_column(Text, nullable=False)
synced_at: Mapped[str | None] = mapped_column(String, nullable=True) synced_at: Mapped[str | None] = mapped_column(String, nullable=True)
mcp_busy: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
class ScheduledTaskRow(Base): class ScheduledTaskRow(Base):
@@ -185,6 +185,7 @@ def create_internal_router(
address=payload.address, address=payload.address,
allow_device_takeover=allow_device_takeover, allow_device_takeover=allow_device_takeover,
planner_transport=payload.planner_transport, planner_transport=payload.planner_transport,
mcp_busy_device_ids=payload.mcp_busy_device_ids,
) )
policy = pool.store.get_host_governance_policy(host_id) policy = pool.store.get_host_governance_policy(host_id)
policy_revision = policy.revision if policy is not None else 0 policy_revision = policy.revision if policy is not None else 0
@@ -0,0 +1,28 @@
"""Add mcp_busy flag column to pooled_devices."""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "0014_pooled_device_mcp_busy"
down_revision = "0013_task_cancellation"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"pooled_devices",
sa.Column(
"mcp_busy",
sa.Boolean(),
nullable=False,
server_default=sa.false(),
),
)
def downgrade() -> None:
op.drop_column("pooled_devices", "mcp_busy")
+12 -1
View File
@@ -47,6 +47,7 @@ class PooledDevice:
status: PooledDeviceStatus status: PooledDeviceStatus
capability_tags: list[str] = field(default_factory=list) capability_tags: list[str] = field(default_factory=list)
synced_at: datetime | None = None synced_at: datetime | None = None
mcp_busy: bool = False
class DevicePool: class DevicePool:
@@ -64,6 +65,7 @@ class DevicePool:
address: str | None = None, address: str | None = None,
planner_transport: Literal["direct", "cloud"] = "direct", planner_transport: Literal["direct", "cloud"] = "direct",
allow_device_takeover: bool = False, allow_device_takeover: bool = False,
mcp_busy_device_ids: list[str] | None = None,
) -> None: ) -> None:
"""Push a host's current device snapshot into the pool. """Push a host's current device snapshot into the pool.
@@ -78,7 +80,13 @@ class DevicePool:
last_seen_at=now, last_seen_at=now,
planner_transport=planner_transport, planner_transport=planner_transport,
) )
devices = [self._to_pooled(device, host_id, now) for device in snapshot] busy_set = set(mcp_busy_device_ids or [])
devices = [
self._to_pooled(
device, host_id, now, mcp_busy=device.id in busy_set
)
for device in snapshot
]
if allow_device_takeover: if allow_device_takeover:
self.store.replace_host_devices( self.store.replace_host_devices(
host_id, host_id,
@@ -119,6 +127,8 @@ class DevicePool:
device: Device, device: Device,
host_id: str, host_id: str,
synced_at: datetime, synced_at: datetime,
*,
mcp_busy: bool = False,
) -> PooledDevice: ) -> PooledDevice:
raw_status = ( raw_status = (
device.status if device.status in _HOST_REPORTED_STATUSES else "idle" device.status if device.status in _HOST_REPORTED_STATUSES else "idle"
@@ -131,6 +141,7 @@ class DevicePool:
status=raw_status, # type: ignore[arg-type] status=raw_status, # type: ignore[arg-type]
capability_tags=tags, capability_tags=tags,
synced_at=synced_at, synced_at=synced_at,
mcp_busy=mcp_busy,
) )
def _is_stale(self, host: HostRegistration, now: datetime) -> bool: def _is_stale(self, host: HostRegistration, now: datetime) -> bool:
@@ -204,6 +204,8 @@ class TaskScheduler:
def _matches(device: "PooledDevice", constraints: TaskConstraints) -> bool: def _matches(device: "PooledDevice", constraints: TaskConstraints) -> bool:
if device.mcp_busy:
return False
if constraints.target_host_id and device.host_id != constraints.target_host_id: if constraints.target_host_id and device.host_id != constraints.target_host_id:
return False return False
if ( if (
@@ -342,6 +342,7 @@ class SQLAlchemyCloudRepository:
ensure_ascii=False, ensure_ascii=False,
), ),
synced_at=_iso(device.synced_at) if device.synced_at else None, synced_at=_iso(device.synced_at) if device.synced_at else None,
mcp_busy=getattr(device, "mcp_busy", False),
) )
for device in devices for device in devices
] ]
@@ -2232,6 +2233,7 @@ def _device_from_row(row: PooledDeviceRow) -> Any:
status=row.status, status=row.status,
capability_tags=tags, capability_tags=tags,
synced_at=_parse_dt(row.synced_at), synced_at=_parse_dt(row.synced_at),
mcp_busy=bool(getattr(row, "mcp_busy", False)),
) )
@@ -0,0 +1,72 @@
"""Tests for the cloud DevicePool, focused on the mcp_busy flag plumbing."""
from __future__ import annotations
import tempfile
from pathlib import Path
import pytest
from cloud.config import CloudConfig
from cloud.pool import DevicePool
from cloud.store import CloudStore
from core.models import Device
@pytest.fixture
def pool() -> DevicePool:
"""Build a fresh DevicePool backed by a temporary SQLite file."""
with tempfile.TemporaryDirectory() as tmp:
store = CloudStore(Path(tmp) / "cloud.sqlite3")
try:
yield DevicePool(store=store, config=CloudConfig())
finally:
store.close()
def _device(device_id: str, *, status: str = "idle") -> Device:
return Device(
id=device_id,
driver_type="wda",
status=status, # type: ignore[arg-type]
capability_tags=[],
)
def test_sync_host_devices_marks_mcp_busy_devices(pool: DevicePool) -> None:
"""When a host reports device-1 as MCP-busy, the pool PooledDevice for
device-1 has mcp_busy=True."""
pool.sync_host_devices(
"host-1",
[_device("device-1", status="idle")],
mcp_busy_device_ids=["device-1"],
)
devices = pool.list_devices()
busy = [d for d in devices if d.device_id == "device-1"]
assert len(busy) == 1
assert busy[0].mcp_busy is True
def test_sync_host_devices_default_mcp_busy_is_false(pool: DevicePool) -> None:
pool.sync_host_devices(
"host-1", [_device("device-1", status="idle")]
)
devices = pool.list_devices()
assert devices[0].mcp_busy is False
def test_sync_host_devices_clears_mcp_busy_on_next_sync(
pool: DevicePool,
) -> None:
"""MCP releases device -> next heartbeat without device in
mcp_busy_device_ids -> pool reflects mcp_busy=False."""
pool.sync_host_devices(
"host-1",
[_device("device-1", status="idle")],
mcp_busy_device_ids=["device-1"],
)
pool.sync_host_devices(
"host-1", [_device("device-1", status="idle")]
)
devices = pool.list_devices()
assert devices[0].mcp_busy is False
@@ -0,0 +1,53 @@
"""Tests for the cloud TaskScheduler, focused on skipping MCP-busy devices."""
from __future__ import annotations
import tempfile
from pathlib import Path
import pytest
from cloud.config import CloudConfig
from cloud.pool import DevicePool
from cloud.scheduler import TaskConstraints, TaskScheduler
from cloud.store import CloudStore
from core.models import Device
@pytest.fixture
def pool() -> DevicePool:
with tempfile.TemporaryDirectory() as tmp:
store = CloudStore(Path(tmp) / "cloud.sqlite3")
try:
yield DevicePool(store=store, config=CloudConfig())
finally:
store.close()
def _device(device_id: str, *, status: str = "idle") -> Device:
return Device(
id=device_id,
driver_type="wda",
status=status, # type: ignore[arg-type]
capability_tags=[],
)
def test_mcp_busy_device_is_skipped_by_scheduler(pool: DevicePool) -> None:
"""A device with mcp_busy=True is not selected for assignment.
Two devices exist: dev-busy (idle status, mcp_busy=True) and dev-idle
(idle status, mcp_busy=False). One task is submitted with no
constraints, so both are candidates before the mcp_busy filter.
The scheduler must pick dev-idle.
"""
pool.sync_host_devices(
"host-1",
[_device("dev-busy"), _device("dev-idle")],
mcp_busy_device_ids=["dev-busy"],
)
scheduler = TaskScheduler(pool=pool, store=pool.store, config=CloudConfig())
scheduler.submit(goal="test", constraints=TaskConstraints())
assignments = scheduler.assign()
assert len(assignments) == 1
assert assignments[0].device_id == "dev-idle"