dynamic type registry

This commit is contained in:
Sami Khan
2026-01-22 11:36:50 +05:00
parent a9db83ba6b
commit 1ea358b808
21 changed files with 293 additions and 184 deletions
+6 -5
View File
@@ -46,8 +46,8 @@ from exo.shared.types.api import (
)
from exo.shared.types.chunks import TokenChunk
from exo.shared.types.commands import (
BaseCommand,
ChatCompletion,
Command,
CreateInstance,
DeleteInstance,
ForwarderCommand,
@@ -66,6 +66,7 @@ from exo.shared.types.models import ModelId, ModelMetadata
from exo.shared.types.state import State
from exo.shared.types.tasks import ChatCompletionTaskParams
from exo.shared.types.worker.instances import (
BaseInstance,
Instance,
InstanceId,
InstanceMeta,
@@ -317,7 +318,7 @@ class API:
sharding: Sharding = Sharding.Pipeline,
instance_meta: InstanceMeta = InstanceMeta.MlxRing,
min_nodes: int = 1,
) -> Instance:
) -> BaseInstance:
model_meta = await resolve_model_meta(model_id)
try:
@@ -449,7 +450,7 @@ class API:
model_id=card.model_id,
sharding=sharding,
instance_meta=instance_meta,
instance=instance,
instance=cast(Instance, instance),
memory_delta_by_node=memory_delta_by_node or None,
error=None,
)
@@ -458,7 +459,7 @@ class API:
return PlacementPreviewResponse(previews=previews)
def get_instance(self, instance_id: InstanceId) -> Instance:
def get_instance(self, instance_id: InstanceId) -> BaseInstance:
if instance_id not in self.state.instances:
raise HTTPException(status_code=404, detail="Instance not found")
return self.state.instances[instance_id]
@@ -808,7 +809,7 @@ class API:
if message.clock > self.last_completed_election:
self.paused = True
async def _send(self, command: Command):
async def _send(self, command: BaseCommand):
while self.paused:
await self.paused_ev.wait()
await self.command_sender.send(
+2 -1
View File
@@ -216,7 +216,8 @@ class Master:
IndexedEvent(idx=i, event=self._event_log[i])
)
case _:
# Plugin-managed commands are handled above
# Plugin commands should be handled by registry above;
# this is a safety fallback for unhandled commands
pass
for event in generated_events:
await self.event_sender.send(event)
+9 -14
View File
@@ -26,7 +26,7 @@ from exo.shared.types.memory import Memory
from exo.shared.types.models import ModelId
from exo.shared.types.profiling import MemoryUsage, NodeNetworkInfo
from exo.shared.types.worker.instances import (
Instance,
BaseInstance,
InstanceId,
InstanceMeta,
MlxJacclInstance,
@@ -43,8 +43,8 @@ def random_ephemeral_port() -> int:
def add_instance_to_placements(
command: CreateInstance,
topology: Topology,
current_instances: Mapping[InstanceId, Instance],
) -> Mapping[InstanceId, Instance]:
current_instances: Mapping[InstanceId, BaseInstance],
) -> Mapping[InstanceId, BaseInstance]:
# TODO: validate against topology
return {**current_instances, command.instance.instance_id: command.instance}
@@ -53,10 +53,10 @@ def add_instance_to_placements(
def place_instance(
command: PlaceInstance,
topology: Topology,
current_instances: Mapping[InstanceId, Instance],
current_instances: Mapping[InstanceId, BaseInstance],
node_memory: Mapping[NodeId, MemoryUsage],
node_network: Mapping[NodeId, NodeNetworkInfo],
) -> dict[InstanceId, Instance]:
) -> dict[InstanceId, BaseInstance]:
cycles = topology.get_cycles()
candidate_cycles = list(filter(lambda it: len(it) >= command.min_nodes, cycles))
cycles_with_sufficient_memory = filter_cycles_by_memory(
@@ -159,19 +159,14 @@ def place_instance(
hosts_by_node=hosts_by_node,
ephemeral_port=ephemeral_port,
)
case _:
# Plugin-managed instance types have their own placement functions
raise ValueError(
f"Instance type {command.instance_meta} must use plugin placement"
)
return target_instances
def delete_instance(
command: DeleteInstance,
current_instances: Mapping[InstanceId, Instance],
) -> dict[InstanceId, Instance]:
current_instances: Mapping[InstanceId, BaseInstance],
) -> dict[InstanceId, BaseInstance]:
target_instances = dict(deepcopy(current_instances))
if command.instance_id in target_instances:
del target_instances[command.instance_id]
@@ -180,8 +175,8 @@ def delete_instance(
def get_transition_events(
current_instances: Mapping[InstanceId, Instance],
target_instances: Mapping[InstanceId, Instance],
current_instances: Mapping[InstanceId, BaseInstance],
target_instances: Mapping[InstanceId, BaseInstance],
) -> Sequence[Event]:
events: list[Event] = []
+10 -24
View File
@@ -4,11 +4,13 @@ This module provides the plugin architecture for extending exo with custom
workload types (simulations, ML frameworks, etc.) without modifying core code.
"""
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from exo.plugins.base import EXOPlugin, PluginCommand, PluginInstance
from exo.plugins.registry import PluginRegistry, discover_plugins
from exo.plugins.base import EXOPlugin, PluginCommand, PluginInstance
from exo.plugins.registry import PluginRegistry, discover_plugins
from exo.plugins.type_registry import (
command_registry,
event_registry,
instance_registry,
)
__all__ = [
"EXOPlugin",
@@ -16,23 +18,7 @@ __all__ = [
"PluginInstance",
"PluginRegistry",
"discover_plugins",
"command_registry",
"event_registry",
"instance_registry",
]
def __getattr__(name: str) -> Any: # pyright: ignore[reportAny]
"""Lazy import to avoid circular dependencies."""
if name in ("EXOPlugin", "PluginCommand", "PluginInstance"):
from exo.plugins.base import EXOPlugin, PluginCommand, PluginInstance
return {
"EXOPlugin": EXOPlugin,
"PluginCommand": PluginCommand,
"PluginInstance": PluginInstance,
}[name]
if name in ("PluginRegistry", "discover_plugins"):
from exo.plugins.registry import PluginRegistry, discover_plugins
return {"PluginRegistry": PluginRegistry, "discover_plugins": discover_plugins}[
name
]
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
+3 -3
View File
@@ -15,7 +15,7 @@ from exo.utils.pydantic_ext import TaggedModel
if TYPE_CHECKING:
from exo.shared.topology import Topology
from exo.shared.types.worker.instances import BoundInstance, Instance
from exo.shared.types.worker.instances import BaseInstance, BoundInstance
from exo.utils.channels import MpReceiver, MpSender
from exo.worker.runner.runner_supervisor import RunnerSupervisor
@@ -113,7 +113,7 @@ class EXOPlugin(ABC):
self,
command: Any, # pyright: ignore[reportAny]
topology: "Topology",
current_instances: Mapping[InstanceId, "Instance"],
current_instances: Mapping[InstanceId, "BaseInstance"],
) -> Sequence[Event]:
"""Process a command and return events to emit.
@@ -140,7 +140,7 @@ class EXOPlugin(ABC):
def plan_task(
self,
runners: Mapping[RunnerId, "RunnerSupervisor"],
instances: Mapping[InstanceId, "Instance"],
instances: Mapping[InstanceId, "BaseInstance"],
) -> Task | None:
"""Plan the next task for plugin instances.
+2 -2
View File
@@ -3,7 +3,7 @@
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from exo.shared.types.commands import Command
from exo.shared.types.commands import BaseCommand
from exo.shared.types.common import NodeId
from exo.shared.types.state import State
@@ -17,5 +17,5 @@ class PluginContext:
"""
state: State
send_command: Callable[[Command], Awaitable[None]]
send_command: Callable[[BaseCommand], Awaitable[None]]
node_id: NodeId
@@ -1,32 +1,15 @@
"""FLASH Plugin - MPI-based simulation support for Exo."""
from typing import TYPE_CHECKING, Any
# Import types directly (these don't cause circular imports)
from exo.plugins.implementations.flash.plugin import FLASHPlugin
from exo.plugins.implementations.flash.types import (
FLASHInstance,
LaunchFLASH,
StopFLASH,
)
if TYPE_CHECKING:
from exo.plugins.implementations.flash.plugin import FLASHPlugin
__all__ = ["FLASHPlugin", "FLASHInstance", "LaunchFLASH", "StopFLASH", "register"]
def register() -> "FLASHPlugin":
def register() -> FLASHPlugin:
"""Entry point for plugin discovery."""
# Lazy import to avoid circular imports during module loading
from exo.plugins.implementations.flash.plugin import FLASHPlugin
return FLASHPlugin()
# For backwards compatibility, allow importing FLASHPlugin from this module
def __getattr__(name: str) -> Any: # pyright: ignore[reportAny]
if name == "FLASHPlugin":
from exo.plugins.implementations.flash.plugin import FLASHPlugin
return FLASHPlugin
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
@@ -11,7 +11,7 @@ from exo.shared.types.common import Host, NodeId
from exo.shared.types.memory import Memory
from exo.shared.types.models import ModelId, ModelMetadata
from exo.shared.types.topology import SocketConnection
from exo.shared.types.worker.instances import Instance, InstanceId
from exo.shared.types.worker.instances import BaseInstance, InstanceId
from exo.shared.types.worker.runners import (
RunnerId,
ShardAssignments,
@@ -22,8 +22,8 @@ from exo.shared.types.worker.shards import PipelineShardMetadata
def place_flash_instance(
command: LaunchFLASH,
topology: Topology,
current_instances: Mapping[InstanceId, Instance],
) -> dict[InstanceId, Instance]:
current_instances: Mapping[InstanceId, BaseInstance],
) -> dict[InstanceId, BaseInstance]:
"""Place a FLASH simulation instance across available nodes.
Unlike MLX instances which use ring/JACCL topology for tensor parallelism,
@@ -31,7 +31,7 @@ def place_flash_instance(
node IPs so the runner can generate an MPI hostfile.
"""
instance_id = InstanceId()
target_instances: dict[InstanceId, Instance] = dict(deepcopy(current_instances))
target_instances: dict[InstanceId, BaseInstance] = dict(deepcopy(current_instances))
all_nodes = list(topology.list_nodes())
@@ -4,14 +4,14 @@ from collections.abc import Mapping
from exo.plugins.implementations.flash.types import FLASHInstance
from exo.shared.types.tasks import LoadModel, Task
from exo.shared.types.worker.instances import Instance, InstanceId
from exo.shared.types.worker.instances import BaseInstance, InstanceId
from exo.shared.types.worker.runners import RunnerId, RunnerIdle
from exo.worker.runner.runner_supervisor import RunnerSupervisor
def plan_flash(
runners: Mapping[RunnerId, RunnerSupervisor],
instances: Mapping[InstanceId, Instance],
instances: Mapping[InstanceId, BaseInstance],
) -> Task | None:
"""Plan tasks specifically for FLASH instances.
@@ -21,7 +21,7 @@ from exo.shared.topology import Topology
from exo.shared.types.commands import DeleteInstance
from exo.shared.types.events import Event
from exo.shared.types.tasks import Task
from exo.shared.types.worker.instances import BoundInstance, Instance, InstanceId
from exo.shared.types.worker.instances import BaseInstance, BoundInstance, InstanceId
from exo.shared.types.worker.runners import RunnerId
from exo.utils.channels import MpReceiver, MpSender
from exo.worker.runner.runner_supervisor import RunnerSupervisor
@@ -60,7 +60,7 @@ class FLASHPlugin(EXOPlugin):
self,
command: Any, # pyright: ignore[reportAny]
topology: Topology,
current_instances: Mapping[InstanceId, Instance],
current_instances: Mapping[InstanceId, BaseInstance],
) -> Sequence[Event]:
from exo.master.placement import delete_instance, get_transition_events
@@ -81,7 +81,7 @@ class FLASHPlugin(EXOPlugin):
def plan_task(
self,
runners: Mapping[RunnerId, RunnerSupervisor],
instances: Mapping[InstanceId, Instance],
instances: Mapping[InstanceId, BaseInstance],
) -> Task | None:
return plan_flash(runners, instances)
+14 -43
View File
@@ -1,34 +1,21 @@
"""FLASH plugin types - commands and instances."""
# ruff: noqa: I001 - Import order intentional for Pydantic model_rebuild
from __future__ import annotations
from typing import TYPE_CHECKING
from pydantic import Field
from exo.shared.types.common import CommandId, Host, NodeId
from exo.shared.types.worker.runners import ShardAssignments
from exo.utils.pydantic_ext import TaggedModel
if TYPE_CHECKING:
from exo.shared.types.worker.instances import InstanceId
from exo.shared.types.worker.runners import RunnerId
from exo.shared.types.worker.shards import (
PipelineShardMetadata,
TensorShardMetadata,
)
from exo.plugins.type_registry import command_registry, instance_registry
from exo.shared.types.commands import BaseCommand
from exo.shared.types.common import Host, NodeId
from exo.shared.types.worker.instances import BaseInstance, InstanceId
from exo.shared.types.worker.runners import RunnerId
from exo.shared.types.worker.shards import ShardMetadata
# ============================================================================
# Commands
# ============================================================================
class LaunchFLASH(TaggedModel):
@command_registry.register
class LaunchFLASH(BaseCommand):
"""Command to launch a FLASH MPI simulation."""
command_id: CommandId = Field(default_factory=CommandId)
simulation_name: str
flash_executable_path: str
parameter_file_path: str
@@ -40,11 +27,11 @@ class LaunchFLASH(TaggedModel):
hosts: str = ""
class StopFLASH(TaggedModel):
@command_registry.register
class StopFLASH(BaseCommand):
"""Command to stop a running FLASH simulation."""
command_id: CommandId = Field(default_factory=CommandId)
instance_id: "InstanceId"
instance_id: InstanceId
# ============================================================================
@@ -52,7 +39,8 @@ class StopFLASH(TaggedModel):
# ============================================================================
class FLASHInstance(TaggedModel):
@instance_registry.register
class FLASHInstance(BaseInstance):
"""Instance for FLASH MPI simulation.
Unlike MLX instances which do tensor parallelism, FLASH instances
@@ -60,8 +48,6 @@ class FLASHInstance(TaggedModel):
MPI ranks of the FLASH simulation.
"""
instance_id: "InstanceId"
shard_assignments: ShardAssignments
hosts_by_node: dict[NodeId, list[Host]]
flash_executable_path: str
parameter_file_path: str
@@ -72,20 +58,5 @@ class FLASHInstance(TaggedModel):
coordinator_ip: str
network_interface: str = "en0" # Network interface for MPI (e.g., en0, eth0)
def shard(
self, runner_id: "RunnerId"
) -> "PipelineShardMetadata | TensorShardMetadata | None":
def shard(self, runner_id: RunnerId) -> ShardMetadata | None:
return self.shard_assignments.runner_to_shard.get(runner_id, None)
# Import types into module namespace for Pydantic model_rebuild() to resolve forward refs
from exo.shared.types.worker.instances import InstanceId as InstanceId # noqa: E402, I001
from exo.shared.types.worker.runners import RunnerId as RunnerId # noqa: E402, I001
from exo.shared.types.worker.shards import ( # noqa: E402, I001
PipelineShardMetadata as PipelineShardMetadata,
TensorShardMetadata as TensorShardMetadata,
)
# Rebuild models to resolve forward references
StopFLASH.model_rebuild()
FLASHInstance.model_rebuild()
+84
View File
@@ -0,0 +1,84 @@
"""Dynamic type registry for plugin types.
This module provides a registry system that allows plugins to register their
command and instance types dynamically, eliminating the need for static union
types and avoiding circular imports.
"""
from typing import TypeVar
from loguru import logger
from exo.utils.pydantic_ext import CamelCaseModel
# TypeVar for preserving exact types through the register decorator
_TCls = TypeVar("_TCls", bound=type[CamelCaseModel])
class TypeRegistry[T: CamelCaseModel]:
"""Registry for dynamically registered Pydantic types.
Enables plugins to register their types at import time. Deserialization
uses the class name from the tagged JSON format to look up the correct type.
"""
def __init__(self, name: str) -> None:
self._name = name
self._types: dict[str, type[T]] = {}
def register(self, cls: _TCls) -> _TCls:
"""Decorator to register a type with this registry.
Preserves the exact type through the decorator for proper type checking.
"""
self._types[cls.__name__] = cls # type: ignore[assignment]
logger.debug(f"{self._name}: registered {cls.__name__}")
return cls
def get(self, name: str) -> type[T] | None:
"""Look up a type by class name."""
return self._types.get(name)
def all_types(self) -> dict[str, type[T]]:
"""Return all registered types."""
return dict(self._types)
def deserialize(self, data: dict[str, dict[str, object]] | CamelCaseModel) -> T:
"""Deserialize dict to the appropriate registered type.
Supports two formats:
1. Tagged format: {"ClassName": {...fields...}} - used for network serialization
2. Flat format: {...fields...} - used for API requests, tries each type
"""
# If already deserialized (e.g., from Pydantic), return as-is
if isinstance(data, CamelCaseModel):
return data # type: ignore[return-value]
# Check for tagged format: single key that matches a registered type
if len(data) == 1:
class_name: str = next(iter(data.keys()))
cls = self._types.get(class_name)
if cls is not None:
return cls.model_validate(data[class_name], strict=False)
# Flat format: try each registered type, use first that validates
errors: list[str] = []
for type_name, cls in self._types.items():
try:
return cls.model_validate(data, strict=False)
except Exception as e: # noqa: BLE001
errors.append(f"{type_name}: {e}")
# None matched - provide helpful error
available = ", ".join(self._types.keys())
raise ValueError(
f"{self._name}: could not deserialize data. "
f"Available types: {available}. Errors: {'; '.join(errors[:3])}"
)
# Global registries for commands, instances, events, and tasks
command_registry: TypeRegistry[CamelCaseModel] = TypeRegistry("CommandRegistry")
instance_registry: TypeRegistry[CamelCaseModel] = TypeRegistry("InstanceRegistry")
event_registry: TypeRegistry[CamelCaseModel] = TypeRegistry("EventRegistry")
task_registry: TypeRegistry[CamelCaseModel] = TypeRegistry("TaskRegistry")
+1 -1
View File
@@ -30,7 +30,7 @@ class TypedTopic[T: CamelCaseModel]:
@staticmethod
def serialize(t: T) -> bytes:
return t.model_dump_json().encode("utf-8")
return t.model_dump_json(by_alias=True, serialize_as_any=True).encode("utf-8")
def deserialize(self, b: bytes) -> T:
return self.model_type.model_validate_json(b.decode("utf-8"))
+7 -3
View File
@@ -34,7 +34,7 @@ from exo.shared.types.state import State
from exo.shared.types.tasks import Task, TaskId, TaskStatus
from exo.shared.types.topology import Connection, RDMAConnection
from exo.shared.types.worker.downloads import DownloadProgress
from exo.shared.types.worker.instances import Instance, InstanceId
from exo.shared.types.worker.instances import BaseInstance, InstanceId
from exo.shared.types.worker.runners import RunnerId, RunnerStatus
from exo.utils.info_gatherer.info_gatherer import (
MacmonMetrics,
@@ -81,6 +81,10 @@ def event_apply(event: Event, state: State) -> State:
return apply_topology_edge_created(event, state)
case TopologyEdgeDeleted():
return apply_topology_edge_deleted(event, state)
case _:
# Unknown event types from plugins are ignored
logger.debug(f"Ignoring unknown event type: {type(event).__name__}")
return state
def apply(state: State, event: IndexedEvent) -> State:
@@ -163,7 +167,7 @@ def apply_task_failed(event: TaskFailed, state: State) -> State:
def apply_instance_created(event: InstanceCreated, state: State) -> State:
instance = event.instance
new_instances: Mapping[InstanceId, Instance] = {
new_instances: Mapping[InstanceId, BaseInstance] = {
**state.instances,
instance.instance_id: instance,
}
@@ -171,7 +175,7 @@ def apply_instance_created(event: InstanceCreated, state: State) -> State:
def apply_instance_deleted(event: InstanceDeleted, state: State) -> State:
new_instances: Mapping[InstanceId, Instance] = {
new_instances: Mapping[InstanceId, BaseInstance] = {
iid: inst for iid, inst in state.instances.items() if iid != event.instance_id
}
return state.model_copy(update={"instances": new_instances})
+14 -2
View File
@@ -1,13 +1,19 @@
import time
from typing import Any, Literal
from typing import Any, Literal, cast
from pydantic import BaseModel, Field, field_validator
from pydantic_core import PydanticUseDefault
from exo.plugins.type_registry import instance_registry
from exo.shared.types.common import CommandId
from exo.shared.types.memory import Memory
from exo.shared.types.models import ModelId, ModelMetadata
from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta
from exo.shared.types.worker.instances import (
BaseInstance,
Instance,
InstanceId,
InstanceMeta,
)
from exo.shared.types.worker.shards import Sharding
FinishReason = Literal[
@@ -184,6 +190,12 @@ class PlaceInstanceParams(BaseModel):
class CreateInstanceParams(BaseModel):
instance: Instance
@field_validator("instance", mode="before")
@classmethod
def validate_instance(cls, v: Any) -> BaseInstance: # noqa: ANN401 # pyright: ignore[reportAny]
"""Validate instance using registry to handle both tagged and flat formats."""
return cast(BaseInstance, instance_registry.deserialize(v)) # pyright: ignore[reportAny]
class PlacementPreview(BaseModel):
model_id: ModelId
+30 -8
View File
@@ -1,6 +1,14 @@
# ruff: noqa: I001 - Import order intentional to avoid circular imports
from pydantic import Field
"""Command types for exo.
Commands are registered dynamically via the command_registry, allowing plugins
to add their own command types without modifying this file.
"""
from typing import Any, cast
from pydantic import Field, field_validator
from exo.plugins.type_registry import command_registry
from exo.shared.types.api import ChatCompletionTaskParams
from exo.shared.types.common import CommandId, NodeId
from exo.shared.types.models import ModelMetadata
@@ -8,22 +16,24 @@ from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta
from exo.shared.types.worker.shards import Sharding
from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
# Import FLASH commands from plugin (for serialization compatibility)
from exo.plugins.implementations.flash.types import LaunchFLASH, StopFLASH # noqa: E402, I001
class BaseCommand(TaggedModel):
"""Base class for all commands."""
command_id: CommandId = Field(default_factory=CommandId)
@command_registry.register
class TestCommand(BaseCommand):
__test__ = False
@command_registry.register
class ChatCompletion(BaseCommand):
request_params: ChatCompletionTaskParams
@command_registry.register
class PlaceInstance(BaseCommand):
model_meta: ModelMetadata
sharding: Sharding
@@ -31,22 +41,27 @@ class PlaceInstance(BaseCommand):
min_nodes: int
@command_registry.register
class CreateInstance(BaseCommand):
instance: Instance
@command_registry.register
class DeleteInstance(BaseCommand):
instance_id: InstanceId
@command_registry.register
class TaskFinished(BaseCommand):
finished_command_id: CommandId
@command_registry.register
class RequestEventLog(BaseCommand):
since_idx: int
# Union type for core commands - used by ForwarderCommand for network deserialization
Command = (
TestCommand
| RequestEventLog
@@ -54,12 +69,19 @@ Command = (
| PlaceInstance
| CreateInstance
| DeleteInstance
| LaunchFLASH
| StopFLASH
| TaskFinished
)
class ForwarderCommand(CamelCaseModel):
"""Wrapper for commands that includes origin node."""
origin: NodeId
command: Command
command: BaseCommand
@field_validator("command", mode="before")
@classmethod
def validate_command(cls, v: Any) -> BaseCommand: # noqa: ANN401 # pyright: ignore[reportAny]
"""Validate command, using registry for plugin commands not in Command union."""
# First try the registry (handles both core and plugin commands)
return cast(BaseCommand, command_registry.deserialize(v)) # pyright: ignore[reportAny]
+48 -25
View File
@@ -1,13 +1,15 @@
from datetime import datetime
from typing import Any, cast
from pydantic import Field
from pydantic import Field, field_validator
from exo.plugins.type_registry import event_registry, instance_registry, task_registry
from exo.shared.topology import Connection
from exo.shared.types.chunks import GenerationChunk
from exo.shared.types.common import CommandId, Id, NodeId, SessionId
from exo.shared.types.tasks import Task, TaskId, TaskStatus
from exo.shared.types.tasks import BaseTask, TaskId, TaskStatus
from exo.shared.types.worker.downloads import DownloadProgress
from exo.shared.types.worker.instances import Instance, InstanceId
from exo.shared.types.worker.instances import BaseInstance, InstanceId
from exo.shared.types.worker.runners import RunnerId, RunnerStatus
from exo.utils.info_gatherer.info_gatherer import GatheredInfo
from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
@@ -25,36 +27,53 @@ class BaseEvent(TaggedModel):
_master_time_stamp: None | datetime = None
@event_registry.register
class TestEvent(BaseEvent):
__test__ = False
@event_registry.register
class TaskCreated(BaseEvent):
task_id: TaskId
task: Task
task: BaseTask
@field_validator("task", mode="before")
@classmethod
def validate_task(cls, v: Any) -> BaseTask: # noqa: ANN401 # pyright: ignore[reportAny]
return cast(BaseTask, task_registry.deserialize(v)) # pyright: ignore[reportAny]
@event_registry.register
class TaskAcknowledged(BaseEvent):
task_id: TaskId
@event_registry.register
class TaskDeleted(BaseEvent):
task_id: TaskId
@event_registry.register
class TaskStatusUpdated(BaseEvent):
task_id: TaskId
task_status: TaskStatus
@event_registry.register
class TaskFailed(BaseEvent):
task_id: TaskId
error_type: str
error_message: str
@event_registry.register
class InstanceCreated(BaseEvent):
instance: Instance
instance: BaseInstance
@field_validator("instance", mode="before")
@classmethod
def validate_instance(cls, v: Any) -> BaseInstance: # noqa: ANN401 # pyright: ignore[reportAny]
return cast(BaseInstance, instance_registry.deserialize(v)) # pyright: ignore[reportAny]
def __eq__(self, other: object) -> bool:
if isinstance(other, InstanceCreated):
@@ -63,72 +82,71 @@ class InstanceCreated(BaseEvent):
return False
@event_registry.register
class InstanceDeleted(BaseEvent):
instance_id: InstanceId
@event_registry.register
class RunnerStatusUpdated(BaseEvent):
runner_id: RunnerId
runner_status: RunnerStatus
@event_registry.register
class RunnerDeleted(BaseEvent):
runner_id: RunnerId
@event_registry.register
class NodeTimedOut(BaseEvent):
node_id: NodeId
# TODO: bikeshed this name
@event_registry.register
class NodeGatheredInfo(BaseEvent):
node_id: NodeId
when: str # this is a manually cast datetime overrode by the master when the event is indexed, rather than the local time on the device
info: GatheredInfo
@event_registry.register
class NodeDownloadProgress(BaseEvent):
download_progress: DownloadProgress
@event_registry.register
class ChunkGenerated(BaseEvent):
command_id: CommandId
chunk: GenerationChunk
@event_registry.register
class TopologyEdgeCreated(BaseEvent):
conn: Connection
@event_registry.register
class TopologyEdgeDeleted(BaseEvent):
conn: Connection
Event = (
TestEvent
| TaskCreated
| TaskStatusUpdated
| TaskFailed
| TaskDeleted
| TaskAcknowledged
| InstanceCreated
| InstanceDeleted
| RunnerStatusUpdated
| RunnerDeleted
| NodeTimedOut
| NodeGatheredInfo
| NodeDownloadProgress
| ChunkGenerated
| TopologyEdgeCreated
| TopologyEdgeDeleted
)
# Type alias for backward compatibility - use BaseEvent for type hints
# Actual deserialization uses event_registry
Event = BaseEvent
class IndexedEvent(CamelCaseModel):
"""An event indexed by the master, with a globally unique index"""
idx: int = Field(ge=0)
event: Event
event: BaseEvent
@field_validator("event", mode="before")
@classmethod
def validate_event(cls, v: Any) -> BaseEvent: # noqa: ANN401 # pyright: ignore[reportAny]
return cast(BaseEvent, event_registry.deserialize(v)) # pyright: ignore[reportAny]
class ForwarderEvent(CamelCaseModel):
@@ -137,4 +155,9 @@ class ForwarderEvent(CamelCaseModel):
origin_idx: int = Field(ge=0)
origin: NodeId
session: SessionId
event: Event
event: BaseEvent
@field_validator("event", mode="before")
@classmethod
def validate_event(cls, v: Any) -> BaseEvent: # noqa: ANN401 # pyright: ignore[reportAny]
return cast(BaseEvent, event_registry.deserialize(v)) # pyright: ignore[reportAny]
+12 -2
View File
@@ -16,7 +16,7 @@ from exo.shared.types.profiling import (
)
from exo.shared.types.tasks import Task, TaskId
from exo.shared.types.worker.downloads import DownloadProgress
from exo.shared.types.worker.instances import Instance, InstanceId
from exo.shared.types.worker.instances import BaseInstance, InstanceId
from exo.shared.types.worker.runners import RunnerId, RunnerStatus
from exo.utils.pydantic_ext import CamelCaseModel
@@ -37,7 +37,7 @@ class State(CamelCaseModel):
strict=True,
arbitrary_types_allowed=True,
)
instances: Mapping[InstanceId, Instance] = {}
instances: Mapping[InstanceId, BaseInstance] = {}
runners: Mapping[RunnerId, RunnerStatus] = {}
downloads: Mapping[NodeId, Sequence[DownloadProgress]] = {}
tasks: Mapping[TaskId, Task] = {}
@@ -52,6 +52,16 @@ class State(CamelCaseModel):
node_network: Mapping[NodeId, NodeNetworkInfo] = {}
node_thunderbolt: Mapping[NodeId, NodeThunderboltInfo] = {}
@field_serializer("instances", mode="plain")
def _encode_instances(
self, value: Mapping[InstanceId, BaseInstance]
) -> dict[str, Any]:
"""Serialize instances with full subclass fields."""
return {
str(k): v.model_dump(by_alias=True, serialize_as_any=True)
for k, v in value.items()
}
@field_serializer("topology", mode="plain")
def _encode_topology(self, value: Topology) -> TopologySnapshot:
return value.to_snapshot()
+11 -9
View File
@@ -2,6 +2,7 @@ from enum import Enum
from pydantic import Field
from exo.plugins.type_registry import task_registry
from exo.shared.types.api import ChatCompletionTaskParams
from exo.shared.types.common import CommandId, Id
from exo.shared.types.worker.instances import BoundInstance, InstanceId
@@ -28,26 +29,32 @@ class BaseTask(TaggedModel):
instance_id: InstanceId
@task_registry.register
class CreateRunner(BaseTask): # emitted by Worker
bound_instance: BoundInstance
@task_registry.register
class DownloadModel(BaseTask): # emitted by Worker
shard_metadata: ShardMetadata
@task_registry.register
class LoadModel(BaseTask): # emitted by Worker
pass
@task_registry.register
class ConnectToGroup(BaseTask): # emitted by Worker
pass
@task_registry.register
class StartWarmup(BaseTask): # emitted by Worker
pass
@task_registry.register
class ChatCompletion(BaseTask): # emitted by Master
command_id: CommandId
task_params: ChatCompletionTaskParams
@@ -56,16 +63,11 @@ class ChatCompletion(BaseTask): # emitted by Master
error_message: str | None = Field(default=None)
@task_registry.register
class Shutdown(BaseTask): # emitted by Worker
runner_id: RunnerId
Task = (
CreateRunner
| DownloadModel
| ConnectToGroup
| LoadModel
| StartWarmup
| ChatCompletion
| Shutdown
)
# Type alias for backward compatibility - use BaseTask for type hints
# Actual deserialization uses task_registry
Task = BaseTask
+25 -10
View File
@@ -1,8 +1,15 @@
# ruff: noqa: I001 - Import order intentional to avoid circular imports
"""Instance types for exo.
Instances are registered dynamically via the instance_registry, allowing plugins
to add their own instance types without modifying this file.
"""
from enum import Enum
from typing import Any, cast
from pydantic import model_validator
from pydantic import field_validator, model_validator
from exo.plugins.type_registry import instance_registry
from exo.shared.types.common import Host, Id, NodeId
from exo.shared.types.worker.runners import RunnerId, ShardAssignments, ShardMetadata
from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
@@ -15,10 +22,11 @@ class InstanceId(Id):
class InstanceMeta(str, Enum):
MlxRing = "MlxRing"
MlxJaccl = "MlxJaccl"
FLASH = "FLASH"
class BaseInstance(TaggedModel):
"""Base class for all instance types."""
instance_id: InstanceId
shard_assignments: ShardAssignments
@@ -26,29 +34,36 @@ class BaseInstance(TaggedModel):
return self.shard_assignments.runner_to_shard.get(runner_id, None)
@instance_registry.register
class MlxRingInstance(BaseInstance):
hosts_by_node: dict[NodeId, list[Host]]
ephemeral_port: int
@instance_registry.register
class MlxJacclInstance(BaseInstance):
jaccl_devices: list[list[str | None]]
jaccl_coordinators: dict[NodeId, str]
# Import FLASHInstance from plugin (for serialization compatibility)
from exo.plugins.implementations.flash.types import FLASHInstance # noqa: E402, I001
# TODO: Single node instance
Instance = MlxRingInstance | MlxJacclInstance | FLASHInstance
# Union type for Pydantic validation - tries each type in order
# This is used by API endpoints (dashboard) which send flat format
Instance = MlxRingInstance | MlxJacclInstance
class BoundInstance(CamelCaseModel):
instance: Instance
"""An instance bound to a specific runner on a specific node."""
instance: BaseInstance
bound_runner_id: RunnerId
bound_node_id: NodeId
@field_validator("instance", mode="before")
@classmethod
def validate_instance(cls, v: Any) -> BaseInstance: # noqa: ANN401 # pyright: ignore[reportAny]
"""Validate instance using registry to handle both tagged and flat formats."""
return cast(BaseInstance, instance_registry.deserialize(v)) # pyright: ignore[reportAny]
@property
def bound_shard(self) -> ShardMetadata:
shard = self.instance.shard(self.bound_runner_id)
+4 -4
View File
@@ -22,8 +22,8 @@ from exo.shared.types.worker.downloads import (
DownloadProgress,
)
from exo.shared.types.worker.instances import (
BaseInstance,
BoundInstance,
Instance,
InstanceId,
)
from exo.shared.types.worker.runners import (
@@ -50,7 +50,7 @@ def plan(
download_status: Mapping[ModelId, DownloadProgress],
# gdls is not expected to be fresh
global_download_status: Mapping[NodeId, Sequence[DownloadProgress]],
instances: Mapping[InstanceId, Instance],
instances: Mapping[InstanceId, BaseInstance],
all_runners: Mapping[RunnerId, RunnerStatus], # all global
tasks: Mapping[TaskId, Task],
) -> Task | None:
@@ -79,7 +79,7 @@ def plan(
def _kill_runner(
runners: Mapping[RunnerId, RunnerSupervisor],
all_runners: Mapping[RunnerId, RunnerStatus],
instances: Mapping[InstanceId, Instance],
instances: Mapping[InstanceId, BaseInstance],
) -> Shutdown | None:
for runner in runners.values():
runner_id = runner.bound_instance.bound_runner_id
@@ -102,7 +102,7 @@ def _kill_runner(
def _create_runner(
node_id: NodeId,
runners: Mapping[RunnerId, RunnerSupervisor],
instances: Mapping[InstanceId, Instance],
instances: Mapping[InstanceId, BaseInstance],
) -> CreateRunner | None:
for instance in instances.values():
runner_id = instance.shard_assignments.node_to_runner.get(node_id, None)