cloud
This commit is contained in:
@@ -0,0 +1,190 @@
|
||||
"""Task scheduler: bounded queue + pluggable assignment strategy.
|
||||
|
||||
Capability: ``task-scheduler``.
|
||||
|
||||
Mirrors the "string key -> swappable implementation" registry shape already used
|
||||
by ``driver/registry.py``, ``perception/provider.py``, and ``workflow/conditions.py``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Literal, Protocol, runtime_checkable
|
||||
from uuid import uuid4
|
||||
|
||||
from core.models import utc_now
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from cloud.config import CloudConfig
|
||||
from cloud.dispatch import Assignment
|
||||
from cloud.pool import DevicePool, PooledDevice
|
||||
from cloud.store import CloudStore
|
||||
|
||||
|
||||
ScheduledTaskStatus = Literal["queued", "assigned", "dispatched", "done", "failed"]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TaskConstraints:
|
||||
"""Optional device constraints attached to a task submission."""
|
||||
|
||||
driver_type: str | None = None
|
||||
capability_tags: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ScheduledTask:
|
||||
"""A task submitted to the cloud scheduler, awaiting or undergoing assignment."""
|
||||
|
||||
id: str
|
||||
goal: str | None
|
||||
workflow_definition_id: str | None
|
||||
constraints: TaskConstraints
|
||||
status: ScheduledTaskStatus = "queued"
|
||||
assigned_device_id: str | None = None
|
||||
assigned_host_id: str | None = None
|
||||
created_at: datetime = field(default_factory=utc_now)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class AssignmentStrategy(Protocol):
|
||||
"""Selects one device from a pre-filterd list of matching candidates."""
|
||||
|
||||
def select(
|
||||
self,
|
||||
task: ScheduledTask,
|
||||
candidates: "list[PooledDevice]",
|
||||
) -> "PooledDevice | None": ...
|
||||
|
||||
|
||||
class FifoMatchStrategy:
|
||||
"""Default strategy: return the first candidate, or None if empty.
|
||||
|
||||
The caller is responsible for pre-filtering to idle, constraint-matching
|
||||
devices in submission order. This keeps the strategy trivially replaceable.
|
||||
"""
|
||||
|
||||
def select(
|
||||
self,
|
||||
task: ScheduledTask,
|
||||
candidates: "list[PooledDevice]",
|
||||
) -> "PooledDevice | None":
|
||||
if not candidates:
|
||||
return None
|
||||
return candidates[0]
|
||||
|
||||
|
||||
FIFO_MATCH_STRATEGY_NAME = "fifo_match"
|
||||
DEFAULT_STRATEGIES: dict[str, AssignmentStrategy] = {
|
||||
FIFO_MATCH_STRATEGY_NAME: FifoMatchStrategy(),
|
||||
}
|
||||
|
||||
|
||||
class UnknownAssignmentStrategyError(ValueError):
|
||||
"""Raised when the configured default strategy is not in the registry."""
|
||||
|
||||
|
||||
class QueueFullError(RuntimeError):
|
||||
"""Raised when a submission would exceed ``config.max_queue_depth``."""
|
||||
|
||||
|
||||
class TaskSubmissionValidationError(ValueError):
|
||||
"""Raised when a submission has neither a goal nor a workflow_definition_id."""
|
||||
|
||||
|
||||
class TaskScheduler:
|
||||
"""Accepts task submissions and assigns queued tasks to idle devices."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pool: "DevicePool",
|
||||
store: "CloudStore",
|
||||
config: "CloudConfig",
|
||||
*,
|
||||
strategies: dict[str, AssignmentStrategy] | None = None,
|
||||
) -> None:
|
||||
self.pool = pool
|
||||
self.store = store
|
||||
self.config = config
|
||||
self._strategies = dict(strategies) if strategies is not None else dict(DEFAULT_STRATEGIES)
|
||||
if config.default_assignment_strategy not in self._strategies:
|
||||
raise UnknownAssignmentStrategyError(
|
||||
f"unknown assignment strategy {config.default_assignment_strategy!r}; "
|
||||
f"registered strategies: {sorted(self._strategies)}"
|
||||
)
|
||||
|
||||
def submit(
|
||||
self,
|
||||
goal: str | None = None,
|
||||
workflow_definition_id: str | None = None,
|
||||
constraints: TaskConstraints | None = None,
|
||||
) -> str:
|
||||
if not goal and not workflow_definition_id:
|
||||
raise TaskSubmissionValidationError(
|
||||
"a submission must specify either a goal or a workflow_definition_id"
|
||||
)
|
||||
if self.store.count_queued_tasks() >= self.config.max_queue_depth:
|
||||
raise QueueFullError(
|
||||
f"task queue is full ({self.config.max_queue_depth} queued)"
|
||||
)
|
||||
task = ScheduledTask(
|
||||
id=uuid4().hex,
|
||||
goal=goal,
|
||||
workflow_definition_id=workflow_definition_id,
|
||||
constraints=constraints or TaskConstraints(),
|
||||
status="queued",
|
||||
created_at=utc_now(),
|
||||
)
|
||||
self.store.enqueue_task(task)
|
||||
return task.id
|
||||
|
||||
def assign(self) -> "list[Assignment]":
|
||||
from cloud.dispatch import Assignment
|
||||
|
||||
strategy = self._strategies[self.config.default_assignment_strategy]
|
||||
assignments: list[Assignment] = []
|
||||
queued = self.store.list_queued_tasks() # ordered oldest-first
|
||||
if not queued:
|
||||
return assignments
|
||||
|
||||
devices = self.pool.list_devices()
|
||||
assigned_device_ids: set[str] = set()
|
||||
for task in queued:
|
||||
candidates = [
|
||||
device
|
||||
for device in devices
|
||||
if device.device_id not in assigned_device_ids
|
||||
and device.status == "idle"
|
||||
and _matches(device, task.constraints)
|
||||
]
|
||||
selected = strategy.select(task, candidates)
|
||||
if selected is None:
|
||||
continue
|
||||
assigned_device_ids.add(selected.device_id)
|
||||
self.store.update_task(
|
||||
task.id,
|
||||
status="assigned",
|
||||
assigned_device_id=selected.device_id,
|
||||
assigned_host_id=selected.host_id,
|
||||
)
|
||||
assignments.append(
|
||||
Assignment(
|
||||
task_id=task.id,
|
||||
device_id=selected.device_id,
|
||||
host_id=selected.host_id,
|
||||
goal=task.goal,
|
||||
workflow_definition_id=task.workflow_definition_id,
|
||||
)
|
||||
)
|
||||
return assignments
|
||||
|
||||
|
||||
def _matches(device: "PooledDevice", constraints: TaskConstraints) -> bool:
|
||||
if constraints.driver_type and device.driver_type != constraints.driver_type:
|
||||
return False
|
||||
if constraints.capability_tags:
|
||||
device_tags = set(device.capability_tags)
|
||||
if not all(tag in device_tags for tag in constraints.capability_tags):
|
||||
return False
|
||||
return True
|
||||
Reference in New Issue
Block a user