← back to Exo
refactor: A Lot
e1894bc106e955607ebb320b374e4d6f27c7490b · 2025-07-07 20:19:08 +0100 · Arbion Halili
Files touched
M master/idempotency.pyM shared/types/api.pyM shared/types/events/chunks.pyM shared/types/events/common.pyM shared/types/events/events.pyM shared/types/models/common.pyM shared/types/models/model.pyM shared/types/networking/data_plane.pyM shared/types/profiling/common.pyM shared/types/states/master.pyM shared/types/states/shared.pyM shared/types/states/worker.pyM shared/types/tasks/common.pyM shared/types/worker/commands_runner.pyM shared/types/worker/common.pyM shared/types/worker/downloads.pyM shared/types/worker/instances.pyM shared/types/worker/mlx.pyM shared/types/worker/resource_monitor.pyM shared/types/worker/runners.pyM shared/types/worker/shards.pyM shared/utils.py
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 →