[object Object]

← back to Exo

refactor: A Lot

e1894bc106e955607ebb320b374e4d6f27c7490b · 2025-07-07 20:19:08 +0100 · Arbion Halili

Files touched

Diff

commit e1894bc106e955607ebb320b374e4d6f27c7490b
Author: Arbion Halili <99731180+ToxicPine@users.noreply.github.com>
Date:   Mon Jul 7 20:19:08 2025 +0100

    refactor: A Lot
---
 master/idempotency.py                   |   8 +-
 shared/types/api.py                     |  10 +-
 shared/types/events/chunks.py           |  58 +++++----
 shared/types/events/common.py           | 214 +++++++++++++++++++++-----------
 shared/types/events/events.py           | 152 ++++++++++-------------
 shared/types/models/common.py           |   4 +-
 shared/types/models/model.py            |   4 +-
 shared/types/networking/data_plane.py   |   9 +-
 shared/types/profiling/common.py        |  23 ++--
 shared/types/states/master.py           |  26 ++--
 shared/types/states/shared.py           |   6 +-
 shared/types/states/worker.py           |   4 +-
 shared/types/tasks/common.py            |  15 ++-
 shared/types/worker/commands_runner.py  |  87 +++++++------
 shared/types/worker/common.py           |   3 +-
 shared/types/worker/downloads.py        |   4 +-
 shared/types/worker/instances.py        |   4 +-
 shared/types/worker/mlx.py              |   6 +-
 shared/types/worker/resource_monitor.py |  66 ++++++----
 shared/types/worker/runners.py          |  13 +-
 shared/types/worker/shards.py           |  31 +++--
 shared/utils.py                         |   9 +-
 22 files changed, 422 insertions(+), 334 deletions(-)

diff --git a/master/idempotency.py b/master/idempotency.py
index a761d2ab..508cec6d 100644
--- a/master/idempotency.py
+++ b/master/idempotency.py
@@ -2,19 +2,19 @@ from hashlib import sha3_224 as hasher
 from typing import Sequence, TypeVar
 from uuid import UUID
 
-from shared.types.events.common import EventId, EventTypes, IdemKeyGenerator, State
+from shared.types.events.common import EventCategories, EventId, IdemKeyGenerator, State
 
-EventTypeT = TypeVar("EventTypeT", bound=EventTypes)
+EventCategoryT = TypeVar("EventCategoryT", bound=EventCategories)
 
 
-def get_idem_tag_generator(base: str) -> IdemKeyGenerator[EventTypeT]:
+def get_idem_tag_generator(base: str) -> IdemKeyGenerator[EventCategoryT]:
     """Generates idempotency keys for events.
 
     The keys are generated by hashing the state sequence number against a base string.
     You can pick any base string, **so long as it's not used in any other function that generates idempotency keys**.
     """
 
-    def get_idem_keys(state: State[EventTypeT], num_keys: int) -> Sequence[EventId]:
+    def get_idem_keys(state: State[EventCategoryT], num_keys: int) -> Sequence[EventId]:
         def recurse(n: int, last: bytes) -> Sequence[EventId]:
             if n == 0:
                 return []
diff --git a/shared/types/api.py b/shared/types/api.py
index 1d5d9cfd..f1bdefbf 100644
--- a/shared/types/api.py
+++ b/shared/types/api.py
@@ -1,10 +1,12 @@
 from typing import Literal
-from pydantic import BaseModel
+
 from openai.types.chat.completion_create_params import CompletionCreateParams
+from pydantic import BaseModel
 
 from shared.types.tasks.common import TaskId
 
+
 class ChatTask(BaseModel):
-  task_id: TaskId
-  kind: Literal["chat"] = "chat"
-  task_data: CompletionCreateParams
\ No newline at end of file
+    task_id: TaskId
+    kind: Literal["chat"] = "chat"
+    task_data: CompletionCreateParams
diff --git a/shared/types/events/chunks.py b/shared/types/events/chunks.py
index 67834aca..e75d6e1e 100644
--- a/shared/types/events/chunks.py
+++ b/shared/types/events/chunks.py
@@ -1,53 +1,62 @@
-from typing import Any, Literal, TypeVar, Generic, Annotated
-from collections.abc import AsyncGenerator
 from enum import Enum
+from typing import Annotated, Generic, Literal, TypeVar
+
+from openai.types.chat.chat_completion import ChatCompletion
+from openai.types.chat.chat_completion_chunk import ChatCompletionChunk
 from pydantic import BaseModel, Field, TypeAdapter
 
-from shared.types.tasks.common import TaskId
-from shared.types.models.common import ModelId
 from shared.openai import FinishReason
+from shared.types.models.common import ModelId
+from shared.types.tasks.common import TaskId
+
+OpenAIResponse = (
+    ChatCompletion | ChatCompletionChunk
+)  ## Currently we only support chat completions
+
 
 class ChunkType(str, Enum):
-    token = 'token'
-    image = 'image'
+    token = "token"
+    image = "image"
+
+
+ChunkT = TypeVar("ChunkT", bound=ChunkType)
 
-ChunkT = TypeVar('ChunkT', bound=ChunkType)
 
 class BaseChunk(BaseModel, Generic[ChunkT]):
     task_id: TaskId
     idx: int
     model: ModelId
 
+
 ###
 
+
 class TokenChunkData(BaseModel):
     text: str
     token_id: int
     finish_reason: FinishReason | None = None
 
+
 class ImageChunkData(BaseModel):
     data: bytes
 
+
 ###
 
+
 class TokenChunk(BaseChunk[ChunkType.token]):
     chunk_data: TokenChunkData
-    chunk_type: Literal[ChunkType.token] = Field(
-        default=ChunkType.token, frozen=True
-    )
+    chunk_type: Literal[ChunkType.token] = Field(default=ChunkType.token, frozen=True)
+
 
 class ImageChunk(BaseChunk[ChunkType.image]):
     chunk_data: ImageChunkData
-    chunk_type: Literal[ChunkType.image] = Field(
-        default=ChunkType.image, frozen=True
-    )
+    chunk_type: Literal[ChunkType.image] = Field(default=ChunkType.image, frozen=True)
+
 
 ###
 
-GenerationChunk = Annotated[
-    TokenChunk | ImageChunk,
-    Field(discriminator="chunk_type")
-]
+GenerationChunk = Annotated[TokenChunk | ImageChunk, Field(discriminator="chunk_type")]
 GenerationChunkTypeAdapter: TypeAdapter[GenerationChunk] = TypeAdapter(GenerationChunk)
 
 # my_chunk: dict[str, Any] = TokenChunk(
@@ -64,18 +73,12 @@ GenerationChunkTypeAdapter: TypeAdapter[GenerationChunk] = TypeAdapter(Generatio
 # restored = GenerationChunkTypeAdapter.validate_python(my_chunk)
 # print(restored)
 
-#### OpenAI API Interfaces ### 
-
-from openai.types.chat.chat_completion import ChatCompletion
-from openai.types.chat.chat_completion_chunk import ChatCompletionChunk
-
-OpenAIResponse = ChatCompletion | ChatCompletionChunk ## Currently we only support chat completions
+#### OpenAI API Interfaces ###
 
+"""
 def send_task(task: Any) -> AsyncGenerator[GenerationChunk]:
-    """
-    This is the 'command' - turns the task into an event and pushes to the event queue.
-    Tokens are then read off the event queue and pushed back to the api via an AsyncGenerator.
-    """
+    # This is the 'command' - turns the task into an event and pushes to the event queue.
+    # Tokens are then read off the event queue and pushed back to the api via an AsyncGenerator.
     ...
 
 def parse_chunk_to_openai_response(chunk: GenerationChunk) -> OpenAIResponse:
@@ -87,3 +90,4 @@ async def handle_task(task: Any) -> AsyncGenerator[OpenAIResponse]:
 
     async for chunk in generator:
         yield parse_chunk_to_openai_response(chunk)
+"""
diff --git a/shared/types/events/common.py b/shared/types/events/common.py
index df759c53..6e5f78cf 100644
--- a/shared/types/events/common.py
+++ b/shared/types/events/common.py
@@ -1,4 +1,4 @@
-from enum import Enum
+from enum import Enum, auto
 from typing import (
     Annotated,
     Callable,
@@ -7,8 +7,6 @@ from typing import (
     Sequence,
     Tuple,
     TypeVar,
-    Union,
-    get_args,
 )
 
 from pydantic import BaseModel, Field, TypeAdapter, model_validator
@@ -16,8 +14,12 @@ from pydantic import BaseModel, Field, TypeAdapter, model_validator
 from shared.types.common import NewUUID, NodeId
 
 
-class EventId(NewUUID): pass
-class TimerId(NewUUID): pass
+class EventId(NewUUID):
+    pass
+
+
+class TimerId(NewUUID):
+    pass
 
 
 class MLXEventTypes(str, Enum):
@@ -67,117 +69,186 @@ class TimerEventTypes(str, Enum):
     TimerCreated = "TimerCreated"
     TimerFired = "TimerFired"
 
+
 class ResourceEventTypes(str, Enum):
     ResourceProfiled = "ResourceProfiled"
 
 
-EventTypes = Union[
-    TaskEventTypes,
-    StreamingEventTypes,
-    InstanceEventTypes,
-    InstanceStateEventTypes,
-    NodePerformanceEventTypes,
-    ControlPlaneEventTypes,
-    DataPlaneEventTypes,
-    TimerEventTypes,
-    MLXEventTypes,
-    ResourceEventTypes,
-]
+class EventCategories(str, Enum):
+    TaskEventTypes = auto()
+    StreamingEventTypes = auto()
+    InstanceEventTypes = auto()
+    InstanceStateEventTypes = auto()
+    NodePerformanceEventTypes = auto()
+    ControlPlaneEventTypes = auto()
+    DataPlaneEventTypes = auto()
+    TimerEventTypes = auto()
+    MLXEventTypes = auto()
 
-EventTypeT = TypeVar("EventTypeT", bound=EventTypes)
-TEventType = TypeVar("TEventType", bound=EventTypes, covariant=True)
 
+PossibleEventOfEventTypeT = TypeVar("PossibleEventOfEventTypeT", bound=Enum)
 
-class SecureEventProtocol(Protocol):
-    def check_origin_id(self, origin_id: NodeId) -> bool: ...
+#  T=(A|B) <: U=(A|B|C)  ==>  Event[A|B] <: Event[A|BCategoryOfEventsT_cov = TypeVar(name="CategoryOfEventsT_cov", bound=EventCategories, covariant=True)
+CategoryOfEventsT_cov = TypeVar(
+    name="CategoryOfEventsT_cov", bound=EventCategories, contravariant=True
+)
+CategoryOfEventsT_con = TypeVar(
+    name="CategoryOfEventsT_con", bound=EventCategories, contravariant=True
+)
+CategoryOfEventsT_inv = TypeVar(
+    name="CategoryOfEventsT_inv",
+    bound=EventCategories,
+    covariant=False,
+    contravariant=False,
+)
 
 
-class Event(BaseModel, SecureEventProtocol, Generic[TEventType]):
-    event_type: TEventType
+class Event(BaseModel, Generic[PossibleEventOfEventTypeT]):
+    event_type: PossibleEventOfEventTypeT
+    event_category: EventCategories
     event_id: EventId
 
+    def check_origin_id(self, origin_id: NodeId) -> bool:
+        return True
+
+
+class TaskEvent(Event[TaskEventTypes]):
+    event_type: TaskEventTypes
+
+
+class InstanceEvent(Event[InstanceEventTypes]):
+    event_type: InstanceEventTypes
+
+
+class InstanceStateEvent(Event[InstanceStateEventTypes]):
+    event_type: InstanceStateEventTypes
+
+
+class MLXEvent(Event[MLXEventTypes]):
+    event_type: MLXEventTypes
+
+
+class NodePerformanceEvent(Event[NodePerformanceEventTypes]):
+    event_type: NodePerformanceEventTypes
+
+
+class ControlPlaneEvent(Event[ControlPlaneEventTypes]):
+    event_type: ControlPlaneEventTypes
+
+
+class StreamingEvent(Event[StreamingEventTypes]):
+    event_type: StreamingEventTypes
 
-class WrappedEvent(BaseModel, Generic[TEventType]):
-    event: Event[TEventType]
+
+class DataPlaneEvent(Event[DataPlaneEventTypes]):
+    event_type: DataPlaneEventTypes
+
+
+class TimerEvent(Event[TimerEventTypes]):
+    event_type: TimerEventTypes
+
+
+class ResourceEvent(Event[ResourceEventTypes]):
+    event_type: ResourceEventTypes
+
+
+class WrappedMessage(BaseModel, Generic[PossibleEventOfEventTypeT]):
+    message: Event[PossibleEventOfEventTypeT]
     origin_id: NodeId
 
     @model_validator(mode="after")
-    def check_origin_id(self) -> "WrappedEvent[TEventType]":
-        if self.event.check_origin_id(self.origin_id):
+    def check_origin_id(self) -> "WrappedMessage[PossibleEventOfEventTypeT]":
+        if self.message.check_origin_id(self.origin_id):
             return self
         raise ValueError("Invalid Event: Origin ID Does Not Match")
 
 
-class PersistedEvent(BaseModel, Generic[TEventType]):
-    event: Event[TEventType]
+class PersistedEvent(BaseModel, Generic[PossibleEventOfEventTypeT]):
+    event: Event[PossibleEventOfEventTypeT]
     sequence_number: int = Field(gt=0)
 
 
-class State(BaseModel, Generic[TEventType]):
-    event_types: tuple[TEventType, ...] = get_args(TEventType)
+class State(BaseModel, Generic[CategoryOfEventsT_cov]):
+    event_category: CategoryOfEventsT_cov
     sequence_number: int = Field(default=0, ge=0)
 
 
-AnnotatedEventType = Annotated[Event[EventTypes], Field(discriminator="event_type")]
+AnnotatedEventType = Annotated[
+    Event[EventCategories], Field(discriminator="event_category")
+]
 EventTypeParser: TypeAdapter[AnnotatedEventType] = TypeAdapter(AnnotatedEventType)
 
-Applicator = Callable[[State[EventTypeT], Event[TEventType]], State[EventTypeT]]
-Apply = Callable[[State[EventTypeT], Event[EventTypeT]], State[EventTypeT]]
+
+# it's not possible to enforce this at compile time, so we have to do it at runtime
+def mock_todo[T](something: T | None) -> T: ...
+
+
+def apply(
+    state: State[CategoryOfEventsT_inv], event: Event[CategoryOfEventsT_inv]
+) -> State[CategoryOfEventsT_inv]: ...
+
+
+#  T=(A|B) <: U=(A|B|C)  ==>  Apply[A|B] <: Apply[A|B|C]
 SagaApplicator = Callable[
-    [State[EventTypeT], Event[TEventType]], Sequence[Event[EventTypeT]]
+    [State[CategoryOfEventsT_inv], Event[CategoryOfEventsT_inv]],
+    Sequence[Event[CategoryOfEventsT_inv]],
 ]
-Saga = Callable[[State[EventTypeT], Event[EventTypeT]], Sequence[Event[EventTypeT]]]
-
-StateAndEvent = Tuple[State[EventTypeT], Event[EventTypeT]]
-EffectHandler = Callable[[StateAndEvent[EventTypeT], State[EventTypeT]], None]
-EventPublisher = Callable[[Event[EventTypeT]], None]
+Saga = Callable[
+    [State[CategoryOfEventsT_inv], Event[CategoryOfEventsT_inv]],
+    Sequence[Event[CategoryOfEventsT_inv]],
+]
+Apply = Callable[
+    [State[CategoryOfEventsT_inv], Event[CategoryOfEventsT_inv]],
+    State[CategoryOfEventsT_inv],
+]
+StateAndEvent = Tuple[State[CategoryOfEventsT_inv], Event[CategoryOfEventsT_inv]]
+EffectHandler = Callable[
+    [StateAndEvent[CategoryOfEventsT_inv], State[CategoryOfEventsT_inv]], None
+]
+EventPublisher = Callable[[Event[CategoryOfEventsT_inv]], None]
 
 
-class MutableState(Protocol, Generic[EventTypeT]):
+class MutableState[EventCategoryT: EventCategories](Protocol):
     def apply(
         self,
-        event: Event[TEventType],
-        applicator: Applicator[EventTypeT, TEventType],
-        effect_handlers: Sequence[EffectHandler[TEventType]],
+        event: Event[EventCategoryT],
+        applicator: Apply[EventCategoryT],
+        effect_handlers: Sequence[EffectHandler[EventCategoryT]],
     ) -> None: ...
 
 
 class EventOutbox(Protocol):
-    def send(self, events: Sequence[Event[EventTypeT]]) -> None: ...
+    def send(self, events: Sequence[Event[EventCategories]]) -> None: ...
 
 
-class EventProcessor(Protocol):
-    # TODO: is .update() an anti-pattern?
-    def update(
-        self,
-        state: State[EventTypeT],
-        apply: Apply[EventTypeT],
-        effect_handlers: Sequence[EffectHandler[EventTypeT]],
-    ) -> State[EventTypeT]: ...
-
+#
+#  T=[A|B] <: U=[A|B|C]   =>   EventProcessor[A|B] :> EventProcessor[A|B|C]
+#
+class EventProcessor[EventCategoryT: EventCategories](Protocol):
     def get_events_to_apply(
-        self, state: State[TEventType]
-    ) -> Sequence[Event[TEventType]]: ...
+        self, state: State[EventCategoryT]
+    ) -> Sequence[Event[EventCategoryT]]: ...
 
 
-def get_saga_effect_handler(
-    sagas: Saga[EventTypeT], event_publisher: EventPublisher[EventTypeT]
-) -> EffectHandler[EventTypeT]:
-    def effect_handler(state_and_event: StateAndEvent[EventTypeT]) -> None:
+def get_saga_effect_handler[EventCategoryT: EventCategories](
+    saga: Saga[EventCategoryT], event_publisher: EventPublisher[EventCategoryT]
+) -> EffectHandler[EventCategoryT]:
+    def effect_handler(state_and_event: StateAndEvent[EventCategoryT]) -> None:
         trigger_state, trigger_event = state_and_event
-        for event in sagas(trigger_state, trigger_event):
+        for event in saga(trigger_state, trigger_event):
             event_publisher(event)
 
     return lambda state_and_event, _: effect_handler(state_and_event)
 
 
-def get_effects_from_sagas(
-    sagas: Sequence[Saga[EventTypeT]], event_publisher: EventPublisher[EventTypeT]
-) -> Sequence[EffectHandler[EventTypeT]]:
+def get_effects_from_sagas[EventCategoryT: EventCategories](
+    sagas: Sequence[Saga[EventCategoryT]],
+    event_publisher: EventPublisher[EventCategoryT],
+) -> Sequence[EffectHandler[EventCategoryT]]:
     return [get_saga_effect_handler(saga, event_publisher) for saga in sagas]
 
 
-IdemKeyGenerator = Callable[[State[EventTypeT], int], Sequence[EventId]]
+IdemKeyGenerator = Callable[[State[CategoryOfEventsT_cov], int], Sequence[EventId]]
 
 
 class CommandId(NewUUID):
@@ -190,15 +261,14 @@ class CommandTypes(str, Enum):
     Delete = "Delete"
 
 
-CommandTypeT = TypeVar("CommandTypeT", bound=CommandTypes)
-TCommandType = TypeVar("TCommandType", bound=CommandTypes, covariant=True)
-
-
-class Command(BaseModel, Generic[TEventType, TCommandType]):
-    command_type: TCommandType
+class Command[EventCategoryT: EventCategories, CommandType: CommandTypes](BaseModel):
+    command_type: CommandType
     command_id: CommandId
 
 
+CommandTypeT = TypeVar("CommandTypeT", bound=CommandTypes, covariant=True)
+
 Decide = Callable[
-    [State[EventTypeT], Command[TEventType, TCommandType]], Sequence[Event[EventTypeT]]
+    [State[CategoryOfEventsT_cov], Command[CategoryOfEventsT_cov, CommandTypeT]],
+    Sequence[Event[CategoryOfEventsT_cov]],
 ]
diff --git a/shared/types/events/events.py b/shared/types/events/events.py
index a2c9bc08..1f6422c8 100644
--- a/shared/types/events/events.py
+++ b/shared/types/events/events.py
@@ -5,33 +5,39 @@ from typing import Any, Literal, Tuple
 from pydantic import BaseModel
 
 from shared.types.common import NodeId
-from shared.types.events.common import TimerId
 from shared.types.events.common import (
+    ControlPlaneEvent,
     ControlPlaneEventTypes,
+    DataPlaneEvent,
     DataPlaneEventTypes,
-    Event,
+    InstanceEvent,
     InstanceEventTypes,
+    InstanceStateEvent,
     InstanceStateEventTypes,
+    MLXEvent,
     MLXEventTypes,
+    NodePerformanceEvent,
     NodePerformanceEventTypes,
+    ResourceEvent,
+    ResourceEventTypes,
+    StreamingEvent,
     StreamingEventTypes,
+    TaskEvent,
     TaskEventTypes,
+    TimerEvent,
     TimerEventTypes,
-    ResourceEventTypes,
+    TimerId,
 )
 from shared.types.networking.control_plane import (
     ControlPlaneEdgeId,
     ControlPlaneEdgeType,
 )
 from shared.types.networking.data_plane import (
-    AddressingProtocol,
-    ApplicationProtocol,
     DataPlaneEdge,
     DataPlaneEdgeId,
-    DataPlaneEdgeInfoType,
     DataPlaneEdgeProfile,
 )
-from shared.types.profiling.common import NodePerformanceProfile
+from shared.types.profiling.common import NodePerformanceProfile, ProfiledResourceName
 from shared.types.tasks.common import (
     TaskData,
     TaskId,
@@ -43,167 +49,137 @@ from shared.types.tasks.common import (
 from shared.types.worker.common import InstanceId, NodeStatus
 from shared.types.worker.instances import InstanceData, InstanceStatus
 from shared.types.worker.runners import RunnerId, RunnerState, RunnerStateType
-from shared.types.profiling.common import ProfiledResourceName
 
 
 class TimerData(BaseModel):
     timer_id: TimerId
 
 
-class TaskCreated[TaskTypeT: TaskType](Event[TaskEventTypes.TaskCreated]):
-    event_type: Literal[TaskEventTypes.TaskCreated] = TaskEventTypes.TaskCreated
+class TaskCreated[TaskTypeT: TaskType](TaskEvent):
+    event_type: TaskEventTypes = TaskEventTypes.TaskCreated
     task_id: TaskId
     task_data: TaskData[TaskTypeT]
-    task_state: TaskState[TaskTypeT, Literal[TaskStatusIncompleteType.Pending]]
+    task_state: TaskState[Literal[TaskStatusIncompleteType.Pending], TaskTypeT]
     on_instance: InstanceId
 
 
-class TaskUpdated[TaskTypeT: TaskType](Event[TaskEventTypes.TaskUpdated]):
-    event_type: Literal[TaskEventTypes.TaskUpdated] = TaskEventTypes.TaskUpdated
+class TaskUpdated[TaskTypeT: TaskType](TaskEvent):
+    event_type: TaskEventTypes = TaskEventTypes.TaskUpdated
     task_id: TaskId
-    update_data: TaskState[TaskTypeT, TaskStatusType]
+    update_data: TaskState[TaskStatusType, TaskTypeT]
 
 
-class TaskDeleted(Event[TaskEventTypes.TaskDeleted]):
-    event_type: Literal[TaskEventTypes.TaskDeleted] = TaskEventTypes.TaskDeleted
+class TaskDeleted(TaskEvent):
+    event_type: TaskEventTypes = TaskEventTypes.TaskDeleted
     task_id: TaskId
 
 
-class InstanceCreated(Event[InstanceEventTypes.InstanceCreated]):
-    event_type: Literal[InstanceEventTypes.InstanceCreated] = (
-        InstanceEventTypes.InstanceCreated
-    )
+class InstanceCreated(InstanceEvent):
+    event_type: InstanceEventTypes = InstanceEventTypes.InstanceCreated
     instance_id: InstanceId
     instance_data: InstanceData
     target_status: InstanceStatus
 
 
-class InstanceDeleted(Event[InstanceEventTypes.InstanceDeleted]):
-    event_type: Literal[InstanceEventTypes.InstanceDeleted] = (
-        InstanceEventTypes.InstanceDeleted
-    )
+class InstanceDeleted(InstanceEvent):
+    event_type: InstanceEventTypes = InstanceEventTypes.InstanceDeleted
     instance_id: InstanceId
 
 
-class InstanceStatusUpdated(Event[InstanceEventTypes.InstanceStatusUpdated]):
-    event_type: Literal[InstanceEventTypes.InstanceStatusUpdated] = (
-        InstanceEventTypes.InstanceStatusUpdated
-    )
+class InstanceStatusUpdated(InstanceEvent):
+    event_type: InstanceEventTypes = InstanceEventTypes.InstanceStatusUpdated
     instance_id: InstanceId
     instance_status: InstanceStatus
 
 
-class InstanceRunnerStateUpdated(
-    Event[InstanceStateEventTypes.InstanceRunnerStateUpdated]
-):
-    event_type: Literal[InstanceStateEventTypes.InstanceRunnerStateUpdated] = (
+class InstanceRunnerStateUpdated(InstanceStateEvent):
+    event_type: InstanceStateEventTypes = (
         InstanceStateEventTypes.InstanceRunnerStateUpdated
     )
     instance_id: InstanceId
     state_update: Tuple[RunnerId, RunnerState[RunnerStateType]]
 
 
-class InstanceToBeReplacedAtomically(
-    Event[InstanceEventTypes.InstanceToBeReplacedAtomically]
-):
+class InstanceToBeReplacedAtomically(InstanceEvent):
+    event_type: InstanceEventTypes = InstanceEventTypes.InstanceToBeReplacedAtomically
     transition: Tuple[InstanceId, InstanceId]
 
 
-class InstanceReplacedAtomically(Event[InstanceEventTypes.InstanceReplacedAtomically]):
-    event_type: Literal[InstanceEventTypes.InstanceReplacedAtomically] = (
-        InstanceEventTypes.InstanceReplacedAtomically
-    )
+class InstanceReplacedAtomically(InstanceEvent):
+    event_type: InstanceEventTypes = InstanceEventTypes.InstanceReplacedAtomically
     transition: Tuple[InstanceId, InstanceId]
 
 
-class MLXInferenceSagaPrepare(Event[MLXEventTypes.MLXInferenceSagaPrepare]):
-    event_type: Literal[MLXEventTypes.MLXInferenceSagaPrepare] = (
-        MLXEventTypes.MLXInferenceSagaPrepare
-    )
+class MLXInferenceSagaPrepare(MLXEvent):
+    event_type: MLXEventTypes = MLXEventTypes.MLXInferenceSagaPrepare
     task_id: TaskId
     instance_id: InstanceId
 
 
-class MLXInferenceSagaStartPrepare(Event[MLXEventTypes.MLXInferenceSagaStartPrepare]):
-    event_type: Literal[MLXEventTypes.MLXInferenceSagaStartPrepare] = (
-        MLXEventTypes.MLXInferenceSagaStartPrepare
-    )
+class MLXInferenceSagaStartPrepare(MLXEvent):
+    event_type: MLXEventTypes = MLXEventTypes.MLXInferenceSagaStartPrepare
     task_id: TaskId
     instance_id: InstanceId
 
 
-class NodePerformanceProfiled(Event[NodePerformanceEventTypes.NodePerformanceProfiled]):
-    event_type: Literal[NodePerformanceEventTypes.NodePerformanceProfiled] = (
+class NodePerformanceProfiled(NodePerformanceEvent):
+    event_type: NodePerformanceEventTypes = (
         NodePerformanceEventTypes.NodePerformanceProfiled
     )
     node_id: NodeId
     node_profile: NodePerformanceProfile
 
 
-class WorkerConnected(Event[ControlPlaneEventTypes.WorkerConnected]):
-    event_type: Literal[ControlPlaneEventTypes.WorkerConnected] = (
-        ControlPlaneEventTypes.WorkerConnected
-    )
-    edge: DataPlaneEdge[AddressingProtocol, ApplicationProtocol]
+class WorkerConnected(ControlPlaneEvent):
+    event_type: ControlPlaneEventTypes = ControlPlaneEventTypes.WorkerConnected
+    edge: DataPlaneEdge
 
 
-class WorkerStatusUpdated(Event[ControlPlaneEventTypes.WorkerStatusUpdated]):
-    event_type: Literal[ControlPlaneEventTypes.WorkerStatusUpdated] = (
-        ControlPlaneEventTypes.WorkerStatusUpdated
-    )
+class WorkerStatusUpdated(ControlPlaneEvent):
+    event_type: ControlPlaneEventTypes = ControlPlaneEventTypes.WorkerStatusUpdated
     node_id: NodeId
     node_state: NodeStatus
 
 
-class WorkerDisconnected(Event[ControlPlaneEventTypes.WorkerConnected]):
-    event_type: Literal[ControlPlaneEventTypes.WorkerConnected] = (
-        ControlPlaneEventTypes.WorkerConnected
-    )
+class WorkerDisconnected(ControlPlaneEvent):
+    event_type: ControlPlaneEventTypes = ControlPlaneEventTypes.WorkerConnected
     vertex_id: ControlPlaneEdgeId
 
 
-class ChunkGenerated(Event[StreamingEventTypes.ChunkGenerated]):
-    event_type: Literal[StreamingEventTypes.ChunkGenerated] = (
-        StreamingEventTypes.ChunkGenerated
-    )
+class ChunkGenerated(StreamingEvent):
+    event_type: StreamingEventTypes = StreamingEventTypes.ChunkGenerated
     task_id: TaskId
     instance_id: InstanceId
     chunk: Any
 
 
-class DataPlaneEdgeCreated(Event[DataPlaneEventTypes.DataPlaneEdgeCreated]):
-    event_type: Literal[DataPlaneEventTypes.DataPlaneEdgeCreated] = (
-        DataPlaneEventTypes.DataPlaneEdgeCreated
-    )
+class DataPlaneEdgeCreated(DataPlaneEvent):
+    event_type: DataPlaneEventTypes = DataPlaneEventTypes.DataPlaneEdgeCreated
     vertex: ControlPlaneEdgeType
 
 
-class DataPlaneEdgeProfiled(Event[DataPlaneEventTypes.DataPlaneEdgeProfiled]):
-    event_type: Literal[DataPlaneEventTypes.DataPlaneEdgeProfiled] = (
-        DataPlaneEventTypes.DataPlaneEdgeProfiled
-    )
-    edge_profile: DataPlaneEdgeProfile[Literal[DataPlaneEdgeInfoType.network_profile]]
+class DataPlaneEdgeProfiled(DataPlaneEvent):
+    event_type: DataPlaneEventTypes = DataPlaneEventTypes.DataPlaneEdgeProfiled
+    edge_id: DataPlaneEdgeId
+    edge_profile: DataPlaneEdgeProfile
 
 
-class DataPlaneEdgeDeleted(Event[DataPlaneEventTypes.DataPlaneEdgeDeleted]):
-    event_type: Literal[DataPlaneEventTypes.DataPlaneEdgeDeleted] = (
-        DataPlaneEventTypes.DataPlaneEdgeDeleted
-    )
+class DataPlaneEdgeDeleted(DataPlaneEvent):
+    event_type: DataPlaneEventTypes = DataPlaneEventTypes.DataPlaneEdgeDeleted
     edge_id: DataPlaneEdgeId
 
 
-class TimerScheduled(Event[TimerEventTypes.TimerCreated]):
-    event_type: Literal[TimerEventTypes.TimerCreated] = TimerEventTypes.TimerCreated
+class TimerScheduled(TimerEvent):
+    event_type: TimerEventTypes = TimerEventTypes.TimerCreated
     timer_data: TimerData
 
 
-class TimerFired(Event[TimerEventTypes.TimerFired]):
-    event_type: Literal[TimerEventTypes.TimerFired] = TimerEventTypes.TimerFired
+class TimerFired(TimerEvent):
+    event_type: TimerEventTypes = TimerEventTypes.TimerFired
     timer_data: TimerData
 
-class ResourceProfiled(Event[ResourceEventTypes.ResourceProfiled]):
-    event_type: Literal[ResourceEventTypes.ResourceProfiled] = (
-        ResourceEventTypes.ResourceProfiled
-    ) 
+
+class ResourceProfiled(ResourceEvent):
+    event_type: ResourceEventTypes = ResourceEventTypes.ResourceProfiled
     resource_name: ProfiledResourceName
-    resource_profile: NodePerformanceProfile
\ No newline at end of file
+    resource_profile: NodePerformanceProfile
diff --git a/shared/types/models/common.py b/shared/types/models/common.py
index c65cd884..05e82a34 100644
--- a/shared/types/models/common.py
+++ b/shared/types/models/common.py
@@ -1,3 +1,5 @@
 from shared.types.common import NewUUID
 
-class ModelId(NewUUID): pass
\ No newline at end of file
+
+class ModelId(NewUUID):
+    pass
diff --git a/shared/types/models/model.py b/shared/types/models/model.py
index 8588f043..faa7c3ad 100644
--- a/shared/types/models/model.py
+++ b/shared/types/models/model.py
@@ -1,4 +1,4 @@
-from typing import final, Sequence
+from typing import Sequence, final
 
 from pydantic import BaseModel, TypeAdapter
 
@@ -15,4 +15,4 @@ class ModelInfo(BaseModel):
     model_metadata: ModelMetadata
 
 
-ModelIdAdapter: TypeAdapter[ModelId] = TypeAdapter(ModelId)
\ No newline at end of file
+ModelIdAdapter: TypeAdapter[ModelId] = TypeAdapter(ModelId)
diff --git a/shared/types/networking/data_plane.py b/shared/types/networking/data_plane.py
index acb022eb..9c570973 100644
--- a/shared/types/networking/data_plane.py
+++ b/shared/types/networking/data_plane.py
@@ -3,7 +3,8 @@ from typing import Annotated, Literal, TypeVar, Union, final
 
 from pydantic import BaseModel, Field, IPvAnyAddress, TypeAdapter
 
-from shared.types.common import NewUUID
+from shared.types.common import NewUUID, NodeId
+from shared.types.graphs.common import Edge
 
 
 class DataPlaneEdgeId(NewUUID):
@@ -23,14 +24,14 @@ ApP = TypeVar("ApP", bound=ApplicationProtocol)
 
 
 @final
-class DataPlaneEdgeBenchmarkData(BaseModel):
+class DataPlaneEdgeProfile(BaseModel):
     throughput: float
     latency: float
     jitter: float
 
 
 class CommonDataPlaneEdgeData(BaseModel):
-    edge_data_transfer_rate: DataPlaneEdgeBenchmarkData | None = None
+    edge_data_transfer_rate: DataPlaneEdgeProfile | None = None
 
 
 class MlxEdgeMetadata(BaseModel):
@@ -63,3 +64,5 @@ _DataPlaneEdgeData = Annotated[
     Field(discriminator="addressing_protocol"),
 ]
 DataPlaneEdgeAdapter: TypeAdapter[DataPlaneEdgeData] = TypeAdapter(_DataPlaneEdgeData)
+
+DataPlaneEdge = Edge[DataPlaneEdgeData, DataPlaneEdgeId, NodeId]
diff --git a/shared/types/profiling/common.py b/shared/types/profiling/common.py
index ecf07729..1b318cc7 100644
--- a/shared/types/profiling/common.py
+++ b/shared/types/profiling/common.py
@@ -1,20 +1,22 @@
-from typing import Annotated, Literal, Coroutine, Generic, TypeVar
 from enum import Enum
-from abc import ABC
+from typing import Annotated, Generic, Literal, TypeVar
+
 from pydantic import BaseModel, Field, TypeAdapter
 
 
 class ProfiledResourceName(str, Enum):
-    memory = 'memory'   
-    system = 'system'
+    memory = "memory"
+    system = "system"
+
+
+ProfiledResourceT = TypeVar(name="ProfiledResourceT", bound=ProfiledResourceName)
 
-ProfiledResourceT = TypeVar(name='ProfiledResourceT', bound=ProfiledResourceName)
 
 class BasePerformanceProfile(BaseModel, Generic[ProfiledResourceT]):
     """
     Details a single resource (or resource type) that is being monitored by the resource monitor.
     """
-    pass
+
 
 class MemoryPerformanceProfile(BasePerformanceProfile[ProfiledResourceName.memory]):
     resource_name: Literal[ProfiledResourceName.memory] = Field(
@@ -25,11 +27,13 @@ class MemoryPerformanceProfile(BasePerformanceProfile[ProfiledResourceName.memor
     swap_total: int
     swap_used: int
 
+
 class NetworkInterfaceInfo(BaseModel):
     name: str
     ip_address: str
     type: str
 
+
 class SystemPerformanceProfile(BasePerformanceProfile[ProfiledResourceName.system]):
     resource_name: Literal[ProfiledResourceName.system] = Field(
         default=ProfiledResourceName.system, frozen=True
@@ -39,9 +43,12 @@ class SystemPerformanceProfile(BasePerformanceProfile[ProfiledResourceName.syste
     memory: int
     network_interfaces: list[NetworkInterfaceInfo] = Field(default_factory=list)
 
+
 NodePerformanceProfile = Annotated[
     MemoryPerformanceProfile | SystemPerformanceProfile,
-    Field(discriminator="resource_name")
+    Field(discriminator="resource_name"),
 ]
 
-NodePerformanceProfileTypeAdapter: TypeAdapter[NodePerformanceProfile] = TypeAdapter(NodePerformanceProfile)
\ No newline at end of file
+NodePerformanceProfileTypeAdapter: TypeAdapter[NodePerformanceProfile] = TypeAdapter(
+    NodePerformanceProfile
+)
diff --git a/shared/types/states/master.py b/shared/types/states/master.py
index b6486a86..09a5d584 100644
--- a/shared/types/states/master.py
+++ b/shared/types/states/master.py
@@ -7,17 +7,12 @@ from pydantic import BaseModel
 
 from shared.types.common import NodeId
 from shared.types.events.common import (
-    ControlPlaneEventTypes,
-    DataPlaneEventTypes,
     Event,
-    EventTypes,
-    NodePerformanceEventTypes,
+    EventCategories,
     State,
 )
 from shared.types.graphs.resource_graph import ResourceGraph
 from shared.types.networking.data_plane import (
-    AddressingProtocol,
-    ApplicationProtocol,
     DataPlaneEdge,
     DataPlaneEdgeId,
 )
@@ -46,28 +41,24 @@ class CachePolicy(BaseModel, Generic[CachePolicyTypeT]):
     policy_type: CachePolicyTypeT
 
 
-class NodePerformanceProfileState(State[NodePerformanceEventTypes]):
+class NodePerformanceProfileState(State[EventCategories.NodePerformanceEventTypes]):
     node_profiles: Mapping[NodeId, NodePerformanceProfile]
 
 
-class DataPlaneNetworkState(State[DataPlaneEventTypes]):
+class DataPlaneNetworkState(State[EventCategories.DataPlaneEventTypes]):
     topology: DataPlaneTopology
     history: Sequence[OrphanedPartOfDataPlaneTopology]
 
     def delete_edge(self, edge_id: DataPlaneEdgeId) -> None: ...
-    def add_edge(
-        self, edge: DataPlaneEdge[AddressingProtocol, ApplicationProtocol]
-    ) -> None: ...
+    def add_edge(self, edge: DataPlaneEdge) -> None: ...
 
 
-class ControlPlaneNetworkState(State[ControlPlaneEventTypes]):
+class ControlPlaneNetworkState(State[EventCategories.ControlPlaneEventTypes]):
     topology: ControlPlaneTopology
     history: Sequence[OrphanedPartOfControlPlaneTopology]
 
     def delete_edge(self, edge_id: DataPlaneEdgeId) -> None: ...
-    def add_edge(
-        self, edge: DataPlaneEdge[AddressingProtocol, ApplicationProtocol]
-    ) -> None: ...
+    def add_edge(self, edge: DataPlaneEdge) -> None: ...
 
 
 class MasterState(SharedState):
@@ -87,10 +78,7 @@ def get_inference_plan(
 ) -> Mapping[InstanceId, InstanceData]: ...
 
 
-TransitionEventTypes = EventTypes
-
-
 def get_transition_events(
     current_instances: Mapping[InstanceId, InstanceData],
     target_instances: Mapping[InstanceId, InstanceData],
-) -> Sequence[Event[TransitionEventTypes]]: ...
+) -> Sequence[Event[EventCategories]]: ...
diff --git a/shared/types/states/shared.py b/shared/types/states/shared.py
index 15caa2d0..75e3140e 100644
--- a/shared/types/states/shared.py
+++ b/shared/types/states/shared.py
@@ -4,17 +4,17 @@ from typing import Sequence
 from pydantic import BaseModel
 
 from shared.types.common import NodeId
-from shared.types.events.common import InstanceStateEventTypes, State, TaskEventTypes
+from shared.types.events.common import EventCategories, State
 from shared.types.tasks.common import Task, TaskId, TaskStatusType, TaskType
 from shared.types.worker.common import InstanceId
 from shared.types.worker.instances import BaseInstance
 
 
-class KnownInstances(State[InstanceStateEventTypes]):
+class KnownInstances(State[EventCategories.InstanceStateEventTypes]):
     instances: Mapping[InstanceId, BaseInstance]
 
 
-class Tasks(State[TaskEventTypes]):
+class Tasks(State[EventCategories.TaskEventTypes]):
     tasks: Mapping[TaskId, Task[TaskType, TaskStatusType]]
 
 
diff --git a/shared/types/states/worker.py b/shared/types/states/worker.py
index 02b1fb67..699ecb84 100644
--- a/shared/types/states/worker.py
+++ b/shared/types/states/worker.py
@@ -2,14 +2,14 @@ from collections.abc import Mapping
 
 from shared.types.common import NodeId
 from shared.types.events.common import (
-    ControlPlaneEventTypes,
+    EventCategories,
     State,
 )
 from shared.types.states.shared import SharedState
 from shared.types.worker.common import NodeStatus
 
 
-class NodeStatusState(State[ControlPlaneEventTypes]):
+class NodeStatusState(State[EventCategories.ControlPlaneEventTypes]):
     node_status: Mapping[NodeId, NodeStatus]
 
 
diff --git a/shared/types/tasks/common.py b/shared/types/tasks/common.py
index 886ac51b..7e58c35f 100644
--- a/shared/types/tasks/common.py
+++ b/shared/types/tasks/common.py
@@ -1,6 +1,6 @@
 from collections.abc import Mapping
 from enum import Enum
-from typing import Annotated, Generic, Literal, TypeVar
+from typing import Annotated, Generic, Literal, TypeVar, Union
 
 import openai.types.chat as openai
 from pydantic import BaseModel, Field, TypeAdapter
@@ -51,9 +51,6 @@ class TaskStatusCompleteType(str, Enum):
 TaskStatusType = Union[TaskStatusIncompleteType, TaskStatusCompleteType]
 
 
-TaskStatusTypeT = TypeVar("TaskStatusTypeT", bound=TaskStatusType, covariant=True)
-
-
 class TaskArtifact[TaskTypeT: TaskType, TaskStatusTypeT: TaskStatusType](BaseModel): ...
 
 
@@ -92,15 +89,15 @@ class FailedTaskStatus(TaskStatusUpdate[TaskStatusIncompleteType.Failed]):
     error_message: Mapping[RunnerId, str]
 
 
-class TaskState(BaseModel, Generic[TaskTypeT, TaskStatusTypeT]):
+class TaskState[TaskStatusTypeT: TaskStatusType, TaskTypeT: TaskType](BaseModel):
     task_status: TaskStatusUpdate[TaskStatusTypeT]
     task_artifact: TaskArtifact[TaskTypeT, TaskStatusTypeT]
 
 
-class BaseTask(BaseModel, Generic[TaskTypeT, TaskStatusTypeT]):
+class BaseTask[TaskTypeT: TaskType, TaskStatusTypeT: TaskStatusType](BaseModel):
     task_type: TaskTypeT
     task_data: TaskData[TaskTypeT]
-    task_state: TaskState[TaskTypeT, TaskStatusTypeT]
+    task_state: TaskState[TaskStatusTypeT, TaskTypeT]
     on_instance: InstanceId
 
 
@@ -117,5 +114,7 @@ BaseTaskValidator: TypeAdapter[BaseTask[TaskType, TaskStatusType]] = TypeAdapter
 )
 
 
-class Task(BaseTask[TaskTypeT, TaskStatusTypeT]):
+class Task[TaskTypeT: TaskType, TaskStatusTypeT: TaskStatusType](
+    BaseTask[TaskTypeT, TaskStatusTypeT]
+):
     task_id: TaskId
diff --git a/shared/types/worker/commands_runner.py b/shared/types/worker/commands_runner.py
index 57d66fd7..7f636588 100644
--- a/shared/types/worker/commands_runner.py
+++ b/shared/types/worker/commands_runner.py
@@ -1,91 +1,102 @@
-from typing import Annotated, Generic, Literal, TypeVar
 from enum import Enum
+from typing import Annotated, Generic, Literal, TypeVar
+
 from pydantic import BaseModel, Field, TypeAdapter
 
+from shared.openai import FinishReason
 from shared.types.api import ChatTask
-from shared.types.worker.shards import ShardMeta
 from shared.types.worker.mlx import Host
-from shared.openai import FinishReason
+from shared.types.worker.shards import PartitionStrategy, ShardMetadata
 
 ## Messages passed TO the runner
 
+
 class MessageType(str, Enum):
-    Setup = 'setup'
-    ChatTask = 'chat_task'
-    Exit = 'exit'
+    Setup = "setup"
+    ChatTask = "chat_task"
+    Exit = "exit"
+
+
+MT = TypeVar(name="MT", bound=MessageType)
 
-MT = TypeVar(name='MT', bound=MessageType)
 
 class BaseRunnerMessage(BaseModel, Generic[MT]):
     pass
 
+
 class SetupMessage(BaseRunnerMessage[MessageType.Setup]):
-    type: Literal[MessageType.Setup] = Field(
-        default=MessageType.Setup, frozen=True
-    )
-    model_shard_meta: ShardMeta
+    type: Literal[MessageType.Setup] = Field(default=MessageType.Setup, frozen=True)
+    model_shard_meta: ShardMetadata[PartitionStrategy]
     hosts: list[Host]
 
+
 class ChatTaskMessage(BaseRunnerMessage[MessageType.ChatTask]):
     type: Literal[MessageType.ChatTask] = Field(
         default=MessageType.ChatTask, frozen=True
     )
     task: ChatTask
 
+
 class ExitMessage(BaseRunnerMessage[MessageType.Exit]):
-    type: Literal[MessageType.Exit] = Field(
-        default=MessageType.Exit, frozen=True
-    )
+    type: Literal[MessageType.Exit] = Field(default=MessageType.Exit, frozen=True)
+
 
 RunnerMessage = Annotated[
-    SetupMessage | ChatTaskMessage | ExitMessage,
-    Field(discriminator="type")
+    SetupMessage | ChatTaskMessage | ExitMessage, Field(discriminator="type")
 ]
 RunnerMessageTypeAdapter: TypeAdapter[RunnerMessage] = TypeAdapter(RunnerMessage)
 
 ## Responses passed FROM the runner
 
+
 class RunnerResponseType(str, Enum):
     GenerationResponse = "generation_response"
     FinishedResponse = "finished_response"
     PrintResponse = "print_response"
     ErrorResponse = "error_response"
 
-RRT = TypeVar(name='RRT', bound=RunnerResponseType)
+
+RRT = TypeVar(name="RRT", bound=RunnerResponseType)
+
 
 class BaseRunnerResponse(BaseModel, Generic[RRT]):
     pass
 
+
 class GenerationResponse(BaseRunnerResponse[RunnerResponseType.GenerationResponse]):
-  type: Literal[RunnerResponseType.GenerationResponse] = Field(
-    default=RunnerResponseType.GenerationResponse, frozen=True
-  )
-  text: str
-  token: int
-  # logprobs: Optional[list[float]] = None # too big. we can change to be top-k
-  finish_reason: FinishReason | None = None
+    type: Literal[RunnerResponseType.GenerationResponse] = Field(
+        default=RunnerResponseType.GenerationResponse, frozen=True
+    )
+    text: str
+    token: int
+    # logprobs: Optional[list[float]] = None # too big. we can change to be top-k
+    finish_reason: FinishReason | None = None
+
 
 class PrintResponse(BaseRunnerResponse[RunnerResponseType.PrintResponse]):
-  type: Literal[RunnerResponseType.PrintResponse] = Field(
-    default=RunnerResponseType.PrintResponse, frozen=True
-  )
-  text: str
+    type: Literal[RunnerResponseType.PrintResponse] = Field(
+        default=RunnerResponseType.PrintResponse, frozen=True
+    )
+    text: str
+
 
 class FinishedResponse(BaseRunnerResponse[RunnerResponseType.FinishedResponse]):
-  type: Literal[RunnerResponseType.FinishedResponse] = Field(
-    default=RunnerResponseType.FinishedResponse, frozen=True
-  )
+    type: Literal[RunnerResponseType.FinishedResponse] = Field(
+        default=RunnerResponseType.FinishedResponse, frozen=True
+    )
+
 
 class ErrorResponse(BaseRunnerResponse[RunnerResponseType.ErrorResponse]):
-  type: Literal[RunnerResponseType.ErrorResponse] = Field(
-    default=RunnerResponseType.ErrorResponse, frozen=True
-  )
-  error_type: str
-  error_message: str
-  traceback: str | None = None
+    type: Literal[RunnerResponseType.ErrorResponse] = Field(
+        default=RunnerResponseType.ErrorResponse, frozen=True
+    )
+    error_type: str
+    error_message: str
+    traceback: str | None = None
+
 
 RunnerResponse = Annotated[
     GenerationResponse | PrintResponse | FinishedResponse | ErrorResponse,
-    Field(discriminator="type")
+    Field(discriminator="type"),
 ]
 RunnerResponseTypeAdapter: TypeAdapter[RunnerResponse] = TypeAdapter(RunnerResponse)
diff --git a/shared/types/worker/common.py b/shared/types/worker/common.py
index 786e0e73..5fa78f74 100644
--- a/shared/types/worker/common.py
+++ b/shared/types/worker/common.py
@@ -2,6 +2,7 @@ from enum import Enum
 
 from shared.types.common import NewUUID
 
+
 class InstanceId(NewUUID):
     pass
 
@@ -13,4 +14,4 @@ class RunnerId(NewUUID):
 class NodeStatus(str, Enum):
     Idle = "Idle"
     Running = "Running"
-    Paused = "Paused"
\ No newline at end of file
+    Paused = "Paused"
diff --git a/shared/types/worker/downloads.py b/shared/types/worker/downloads.py
index c539fb9c..c88b2d57 100644
--- a/shared/types/worker/downloads.py
+++ b/shared/types/worker/downloads.py
@@ -15,7 +15,7 @@ from pydantic import BaseModel, Field, PositiveInt
 from shared.types.common import NodeId
 from shared.types.models.common import ModelId
 from shared.types.models.sources import ModelSource
-from shared.types.worker.shards import ShardMeta
+from shared.types.worker.shards import PartitionStrategy, ShardMetadata
 
 
 class DownloadProgressData(BaseModel):
@@ -80,6 +80,6 @@ DownloadEffectHandler = Callable[
 def download_shard(
     model_id: ModelId,
     model_source: ModelSource,
-    shard_meta: ShardMeta,
+    shard_meta: ShardMetadata[PartitionStrategy],
     effect_handlers: Sequence[DownloadEffectHandler],
 ) -> None: ...
diff --git a/shared/types/worker/instances.py b/shared/types/worker/instances.py
index 04884d14..f23b5807 100644
--- a/shared/types/worker/instances.py
+++ b/shared/types/worker/instances.py
@@ -6,9 +6,9 @@ from pydantic import BaseModel
 from shared.types.worker.common import InstanceId
 from shared.types.worker.runners import (
     RunnerId,
-    RunnerPlacement,
     RunnerState,
     RunnerStateType,
+    ShardAssignments,
 )
 
 
@@ -22,7 +22,7 @@ class InstanceState(BaseModel):
 
 
 class InstanceData(BaseModel):
-    runner_placements: RunnerPlacement
+    shard_assignments: ShardAssignments
 
 
 class BaseInstance(BaseModel):
diff --git a/shared/types/worker/mlx.py b/shared/types/worker/mlx.py
index 0d5db1f5..496ef369 100644
--- a/shared/types/worker/mlx.py
+++ b/shared/types/worker/mlx.py
@@ -6,8 +6,8 @@ class Host(BaseModel):
     host: str
     port: int
 
-    @field_validator('port')
-    def check_port(cls, v: int) -> int:
+    @field_validator("port")
+    def check_port(self, v: int) -> int:
         if not (0 <= v <= 65535):
             raise ValueError("Port must be between 0 and 65535")
-        return v
\ No newline at end of file
+        return v
diff --git a/shared/types/worker/resource_monitor.py b/shared/types/worker/resource_monitor.py
index ccb115f3..96eba8d2 100644
--- a/shared/types/worker/resource_monitor.py
+++ b/shared/types/worker/resource_monitor.py
@@ -1,55 +1,73 @@
-from abc import ABC
-from collections.abc import Coroutine
-
 import asyncio
+from abc import ABC, abstractmethod
+from collections.abc import Coroutine
+from typing import Callable, Set
 
 from shared.types.events.events import ResourceProfiled
-from shared.types.profiling.common import NodePerformanceProfile, MemoryPerformanceProfile, SystemPerformanceProfile
+from shared.types.profiling.common import (
+    MemoryPerformanceProfile,
+    NodePerformanceProfile,
+    SystemPerformanceProfile,
+)
+
 
 class EventLog:
-    def append(self, event: ResourceProfiled) -> None:
-        ...
+    def append(self, event: ResourceProfiled) -> None: ...
+
 
 class ResourceCollector(ABC):
     """
     Details a single resource (or resource type) that is being monitored by the resource monitor.
     """
+
     def __init__(self, name: str):
         self.name = name
 
-    async def collect(self) -> NodePerformanceProfile:
-        ...
+    @abstractmethod
+    async def collect(self) -> NodePerformanceProfile: ...
+
 
 class SystemResourceCollector(ResourceCollector):
     def __init__(self):
-        super().__init__('system')
+        super().__init__("system")
+
+    @abstractmethod
+    async def collect(self) -> SystemPerformanceProfile: ...
 
-    async def collect(self) -> SystemPerformanceProfile:
-        ...
 
 class MemoryResourceCollector(ResourceCollector):
     def __init__(self):
-        super().__init__('memory')
+        super().__init__("memory")
+
+    @abstractmethod
+    async def collect(self) -> MemoryPerformanceProfile: ...
 
-    async def collect(self) -> MemoryPerformanceProfile:
-        ...
 
 class ResourceMonitor:
-    def __init__(self, event_outbox: EventLog):
-        self.event_outbox: EventLog = event_outbox
+    def __init__(
+        self,
+        collectors: list[ResourceCollector],
+        effect_handlers: Set[Callable[[NodePerformanceProfile], None]],
+    ):
+        self.effect_handlers: Set[Callable[[NodePerformanceProfile], None]] = (
+            effect_handlers
+        )
+        self.collectors: list[ResourceCollector] = collectors
 
-        self.collectors: list[ResourceCollector] = [
-            SystemResourceCollector(),
-            MemoryResourceCollector(),
-        ]
+        # Since there's no implementation, this breaks the typechecker.
+        # self.collectors: list[ResourceCollector] = [
+        #     SystemResourceCollector(),
+        #     MemoryResourceCollector(),
+        # ]
 
-    async def collect(self) -> list[NodePerformanceProfile]:
+    async def _collect(self) -> list[NodePerformanceProfile]:
         tasks: list[Coroutine[None, None, NodePerformanceProfile]] = [
             collector.collect() for collector in self.collectors
         ]
         return await asyncio.gather(*tasks)
 
-    async def collect_and_publish(self) -> None:
-        profiles = await self.collect()
+    async def collect(self) -> None:
+        profiles = await self._collect()
         for profile in profiles:
-            self.event_outbox.append(profile.to_event())
\ No newline at end of file
+            for effect_handler in self.effect_handlers:
+                effect_handler(profile)
diff --git a/shared/types/worker/runners.py b/shared/types/worker/runners.py
index dca7b290..c7528094 100644
--- a/shared/types/worker/runners.py
+++ b/shared/types/worker/runners.py
@@ -1,6 +1,6 @@
 from collections.abc import Mapping, Sequence
 from enum import Enum
-from typing import Generic, Literal, TypeVar, Self
+from typing import Generic, Literal, TypeVar
 
 from pydantic import BaseModel, model_validator
 
@@ -8,7 +8,7 @@ from shared.types.common import NodeId
 from shared.types.models.common import ModelId
 from shared.types.worker.common import RunnerId
 from shared.types.worker.downloads import BaseDownloadProgress, DownloadStatus
-from shared.types.worker.shards import BaseModelShardMeta, PartitionStrategyT
+from shared.types.worker.shards import PartitionStrategy, ShardMetadata
 
 
 class RunnerStateType(str, Enum):
@@ -55,17 +55,16 @@ class RunnerData(BaseModel):
     )
 
 
-# Runner placement must be consistent in its partitioning strategy across all shards.
-# Using a generic type parameter enforces this constraint at type-checking time.
+PartitionStrategyT = TypeVar(name="PartitionStrategyT", bound=PartitionStrategy)
 
 
-class RunnerPlacement(BaseModel, Generic[PartitionStrategyT]):
+class ShardAssignments(BaseModel):
     model_id: ModelId
-    runner_to_shard: Mapping[RunnerId, BaseModelShardMeta[PartitionStrategyT]]
+    runner_to_shard: Mapping[RunnerId, ShardMetadata[PartitionStrategy]]
     node_to_runner: Mapping[NodeId, Sequence[RunnerId]]
 
     @model_validator(mode="after")
-    def validate_runners_exist(self) -> Self:
+    def validate_runners_exist(self) -> "ShardAssignments":
         for runners in self.node_to_runner.values():
             for runner_id in runners:
                 if runner_id not in self.runner_to_shard:
diff --git a/shared/types/worker/shards.py b/shared/types/worker/shards.py
index 57291a79..5b33457d 100644
--- a/shared/types/worker/shards.py
+++ b/shared/types/worker/shards.py
@@ -1,41 +1,47 @@
 from enum import Enum
-from typing import Generic, TypeVar, Annotated, Literal
+from typing import Annotated, Generic, Literal, TypeVar
 
 from pydantic import BaseModel, DirectoryPath, Field, TypeAdapter
 
 from shared.types.common import NodeId
 from shared.types.models.common import ModelId
 
+
 class PartitionStrategy(str, Enum):
-    pipeline = 'pipeline'
+    pipeline = "pipeline"
+
 
-PartitionStrategyT = TypeVar(name='PartitionStrategyT', bound=PartitionStrategy)
+PartitionStrategyT = TypeVar(name="PartitionStrategyT", bound=PartitionStrategy)
 
-class BaseModelShardMeta(BaseModel, Generic[PartitionStrategyT]):
+
+class ShardMetadata(BaseModel, Generic[PartitionStrategyT]):
     """
     Defines a specific shard of the model that is ready to be run on a device.
     Replaces previous `Shard` object.
     """
+
     device_rank: int
     world_size: int
     model_id: ModelId
-    model_path: DirectoryPath # pydantic DirectoryPath ensures that the directory exists.
+    model_path: DirectoryPath
 
-class PipelineShardMeta(BaseModelShardMeta[PartitionStrategy.pipeline]):
+
+class PipelineShardMeta(ShardMetadata[PartitionStrategy.pipeline]):
     """
     Pipeline parallelism shard meta.
     """
+
     partition_strategy: Literal[PartitionStrategy.pipeline] = Field(
         default=PartitionStrategy.pipeline, frozen=True
     )
     start_layer: Annotated[int, Field(ge=0)]
     end_layer: Annotated[int, Field(ge=0)]
 
-ShardMeta = Annotated[
-    PipelineShardMeta,
-    Field(discriminator="partition_strategy")
-]
-ShardMetaAdapter: TypeAdapter[ShardMeta] = TypeAdapter(ShardMeta)
+
+_ShardMeta = Annotated[PipelineShardMeta, Field(discriminator="partition_strategy")]
+ShardMetaAdapter: TypeAdapter[ShardMetadata[PartitionStrategy]] = TypeAdapter(
+    _ShardMeta
+)
 
 
 class ShardPlacement(BaseModel, Generic[PartitionStrategyT]):
@@ -43,5 +49,6 @@ class ShardPlacement(BaseModel, Generic[PartitionStrategyT]):
     A shard placement is the description of a model distributed across a set of nodes.
     The Generic[PartitionStrategyT] enforces that the shard assignments all use the same partition strategy.
     """
+
     model_id: ModelId
-    shard_assignments: dict[NodeId, BaseModelShardMeta[PartitionStrategyT]]
+    shard_assignments: dict[NodeId, ShardMetadata[PartitionStrategyT]]
diff --git a/shared/utils.py b/shared/utils.py
index da09cb04..bf2be769 100644
--- a/shared/utils.py
+++ b/shared/utils.py
@@ -1,8 +1,9 @@
 from typing import Any, Type, TypeVar
 
-T = TypeVar('T')
+T = TypeVar("T")
 
-def ensure_type(obj: Any, expected_type: Type[T]) -> T: # type: ignore
+
+def ensure_type(obj: Any, expected_type: Type[T]) -> T:  # type: ignore
     if not isinstance(obj, expected_type):
-        raise TypeError(f"Expected {expected_type}, got {type(obj)}") # type: ignore
-    return obj
\ No newline at end of file
+        raise TypeError(f"Expected {expected_type}, got {type(obj)}")  # type: ignore
+    return obj

← 81cf6bce refactor: Simplify networking  ·  back to Exo  ·  fix: Make master hold a queue of task data fe17aaf9 →