← back to Exo
Topology apply
df1fe3af26d0f802770d63787c842b3457265099 · 2025-07-24 14:27:09 +0100 · Gelu Vrabie
Co-authored-by: Gelu Vrabie <gelu@exolabs.net>
Files touched
A shared/apply/__init__.pyR074 shared/types/events/_apply.py shared/apply/apply.pyM shared/types/events/__init__.pyD shared/types/events/_common.pyM shared/types/events/_events.pyM shared/types/events/chunks.pyM shared/types/events/commands.pyM shared/types/events/components.pyM shared/types/topology.py
Diff
commit df1fe3af26d0f802770d63787c842b3457265099
Author: Gelu Vrabie <gelu.vrabie.univ@gmail.com>
Date: Thu Jul 24 14:27:09 2025 +0100
Topology apply
Co-authored-by: Gelu Vrabie <gelu@exolabs.net>
---
shared/apply/__init__.py | 3 +
shared/{types/events/_apply.py => apply/apply.py} | 78 +++-------
shared/types/events/__init__.py | 68 +--------
shared/types/events/_common.py | 87 ------------
shared/types/events/_events.py | 166 +++++++++++++++++++---
shared/types/events/chunks.py | 7 +-
shared/types/events/commands.py | 3 +-
shared/types/events/components.py | 3 +-
shared/types/topology.py | 5 +-
9 files changed, 180 insertions(+), 240 deletions(-)
diff --git a/shared/apply/__init__.py b/shared/apply/__init__.py
new file mode 100644
index 00000000..534e5356
--- /dev/null
+++ b/shared/apply/__init__.py
@@ -0,0 +1,3 @@
+from .apply import apply
+
+__all__ = ["apply"]
\ No newline at end of file
diff --git a/shared/types/events/_apply.py b/shared/apply/apply.py
similarity index 74%
rename from shared/types/events/_apply.py
rename to shared/apply/apply.py
index 205517d9..097a5082 100644
--- a/shared/types/events/_apply.py
+++ b/shared/apply/apply.py
@@ -1,19 +1,13 @@
+import copy
from functools import singledispatch
from typing import Mapping, TypeVar
# from shared.topology import Topology
from shared.types.common import NodeId
-from shared.types.events._events import Event
-from shared.types.events.components import EventFromEventLog
-from shared.types.profiling import NodePerformanceProfile
-from shared.types.state import State
-from shared.types.tasks import Task, TaskId
-from shared.types.worker.common import NodeStatus, RunnerId
-from shared.types.worker.instances import BaseInstance, InstanceId, TypeOfInstance
-from shared.types.worker.runners import RunnerStatus
-
-from ._events import (
+from shared.types.events import (
ChunkGenerated,
+ Event,
+ EventFromEventLog,
InstanceActivated,
InstanceCreated,
InstanceDeactivated,
@@ -29,10 +23,14 @@ from ._events import (
TopologyEdgeCreated,
TopologyEdgeDeleted,
TopologyEdgeReplacedAtomically,
- WorkerConnected,
- WorkerDisconnected,
WorkerStatusUpdated,
)
+from shared.types.profiling import NodePerformanceProfile
+from shared.types.state import State
+from shared.types.tasks import Task, TaskId
+from shared.types.worker.common import NodeStatus, RunnerId
+from shared.types.worker.instances import BaseInstance, InstanceId, TypeOfInstance
+from shared.types.worker.runners import RunnerStatus
S = TypeVar("S", bound=State)
@@ -120,61 +118,23 @@ def apply_worker_status_updated(state: State, event: WorkerStatusUpdated) -> Sta
def apply_chunk_generated(state: State, event: ChunkGenerated) -> State:
return state
-# TODO implemente these
-@event_apply.register
-def apply_worker_connected(state: State, event: WorkerConnected) -> State:
- # source_node_id = event.edge.source_node_id
- # sink_node_id = event.edge.sink_node_id
-
- # new_node_status = dict(state.node_status)
- # if source_node_id not in new_node_status:
- # new_node_status[source_node_id] = NodeStatus.Idle
- # if sink_node_id not in new_node_status:
- # new_node_status[sink_node_id] = NodeStatus.Idle
-
- # new_topology = Topology()
- # new_topology.add_connection(event.edge)
-
- # return state.model_copy(update={"node_status": new_node_status, "topology": new_topology})
- return state
-
-@event_apply.register
-def apply_worker_disconnected(state: State, event: WorkerDisconnected) -> State:
- # new_node_status: Mapping[NodeId, NodeStatus] = {nid: status for nid, status in state.node_status.items() if nid != event.vertex_id}
-
- # new_topology = Topology()
-
- # new_history = list(state.history) + [state.topology]
-
- # return state.model_copy(update={
- # "node_status": new_node_status,
- # "topology": new_topology,
- # "history": new_history
- # })
- return state
-
-
@event_apply.register
def apply_topology_edge_created(state: State, event: TopologyEdgeCreated) -> State:
- # new_topology = Topology()
- # new_topology.add_node(event.vertex, event.vertex.node_id)
- # return state.model_copy(update={"topology": new_topology})
- return state
+ topology = copy.copy(state.topology)
+ topology.add_connection(event.edge)
+ return state.model_copy(update={"topology": topology})
@event_apply.register
def apply_topology_edge_replaced_atomically(state: State, event: TopologyEdgeReplacedAtomically) -> State:
- # new_topology = Topology()
- # new_topology.add_connection(event.edge)
- # updated_connection = event.edge.model_copy(update={"connection_profile": event.edge_profile})
- # new_topology.update_connection_profile(updated_connection)
- # return state.model_copy(update={"topology": new_topology})
- return state
+ topology = copy.copy(state.topology)
+ topology.update_connection_profile(event.edge)
+ return state.model_copy(update={"topology": topology})
@event_apply.register
def apply_topology_edge_deleted(state: State, event: TopologyEdgeDeleted) -> State:
- # new_topology = Topology()
- # return state.model_copy(update={"topology": new_topology})
- return state
+ topology = copy.copy(state.topology)
+ topology.remove_connection(event.edge)
+ return state.model_copy(update={"topology": topology})
@event_apply.register
def apply_mlx_inference_saga_prepare(state: State, event: MLXInferenceSagaPrepare) -> State:
diff --git a/shared/types/events/__init__.py b/shared/types/events/__init__.py
index c3052e88..462d460c 100644
--- a/shared/types/events/__init__.py
+++ b/shared/types/events/__init__.py
@@ -4,74 +4,10 @@
# Note: we are implementing internal details here, so importing private stuff is fine!!!
from pydantic import TypeAdapter
-from shared.types.events.components import EventFromEventLog
-
-from ._apply import Event, apply
-from ._common import *
from ._events import *
+from .components import EventFromEventLog
EventParser: TypeAdapter[Event] = TypeAdapter(Event)
"""Type adaptor to parse :class:`Event`s."""
-__all__ = ["Event", "EventParser", "apply", "EventFromEventLog"]
-
-# Event type consistency check - runs after all imports are complete
-def _check_event_type_consistency():
- import types
- import typing
-
- from shared.constants import get_error_reporting_message
-
- from ._common import _BaseEvent, _EventType # pyright: ignore[reportPrivateUsage]
- from ._events import _Event # pyright: ignore[reportPrivateUsage]
-
- # Grab enum values from members
- member_enum_values = [m for m in _EventType]
-
- # grab enum values from the union => scrape the type annotation
- union_enum_values: list[_EventType] = []
- union_classes = list(typing.get_args(_Event))
- for cls in union_classes: # pyright: ignore[reportAny]
- assert issubclass(cls, object), (
- f"{get_error_reporting_message()}",
- f"The class {cls} is NOT a subclass of {object}."
- )
-
- # ensure the first base parameter is ALWAYS _BaseEvent
- base_cls = list(types.get_original_bases(cls))
- assert len(base_cls) >= 1 and issubclass(base_cls[0], object) \
- and issubclass(base_cls[0], _BaseEvent), (
- f"{get_error_reporting_message()}",
- f"The class {cls} does NOT inherit from {_BaseEvent} {typing.get_origin(base_cls[0])}."
- )
-
- # grab type hints and extract the right values from it
- cls_hints = typing.get_type_hints(cls)
- assert "event_type" in cls_hints and \
- typing.get_origin(cls_hints["event_type"]) is typing.Literal, ( # pyright: ignore[reportAny]
- f"{get_error_reporting_message()}",
- f"The class {cls} is missing a {typing.Literal}-annotated `event_type` field."
- )
-
- # make sure the value is an instance of `_EventType`
- enum_value = list(typing.get_args(cls_hints["event_type"]))
- assert len(enum_value) == 1 and isinstance(enum_value[0], _EventType), (
- f"{get_error_reporting_message()}",
- f"The `event_type` of {cls} has a non-{_EventType} literal-type."
- )
- union_enum_values.append(enum_value[0])
-
- # ensure there is a 1:1 bijection between the two
- for m in member_enum_values:
- assert m in union_enum_values, (
- f"{get_error_reporting_message()}",
- f"There is no event-type registered for {m} in {_Event}."
- )
- union_enum_values.remove(m)
- assert len(union_enum_values) == 0, (
- f"{get_error_reporting_message()}",
- f"The following events have multiple event types defined in {_Event}: {union_enum_values}."
- )
-
-
-_check_event_type_consistency()
+__all__ = ["Event", "EventParser", "EventFromEventLog"]
diff --git a/shared/types/events/_common.py b/shared/types/events/_common.py
deleted file mode 100644
index 0090dd32..00000000
--- a/shared/types/events/_common.py
+++ /dev/null
@@ -1,87 +0,0 @@
-from enum import Enum
-from typing import TYPE_CHECKING
-
-if TYPE_CHECKING:
- pass
-
-from pydantic import BaseModel
-
-from shared.types.common import NewUUID, NodeId
-
-# These are exported for use in other modules
-__all__ = ["EventId", "CommandId", "_EventType", "_BaseEvent"]
-
-class EventId(NewUUID):
- """
- Newtype around `NewUUID`
- """
-
-
-# Event base-class boilerplate (you should basically never touch these)
-# Only very specialised registry or serialisation/deserialization logic might need know about these
-class CommandId(NewUUID):
- """
- Newtype around `NewUUID` for command IDs
- """
-
-
-class _EventType(str, Enum):
- """
- Here are all the unique kinds of events that can be sent over the network.
- """
-
- # Task Saga Events
- MLXInferenceSagaPrepare = "MLXInferenceSagaPrepare"
- MLXInferenceSagaStartPrepare = "MLXInferenceSagaStartPrepare"
-
- # Task Events
- TaskCreated = "TaskCreated"
- TaskStateUpdated = "TaskStateUpdated"
- TaskDeleted = "TaskDeleted"
-
- # Streaming Events
- ChunkGenerated = "ChunkGenerated"
-
- # Instance Events
- InstanceCreated = "InstanceCreated"
- InstanceDeleted = "InstanceDeleted"
- InstanceActivated = "InstanceActivated"
- InstanceDeactivated = "InstanceDeactivated"
- InstanceReplacedAtomically = "InstanceReplacedAtomically"
-
- # Runner Status Events
- RunnerStatusUpdated = "RunnerStatusUpdated"
-
- # Node Performance Events
- NodePerformanceMeasured = "NodePerformanceMeasured"
-
- # Topology Events
- TopologyEdgeCreated = "TopologyEdgeCreated"
- TopologyEdgeReplacedAtomically = "TopologyEdgeReplacedAtomically"
- TopologyEdgeDeleted = "TopologyEdgeDeleted"
- WorkerConnected = "WorkerConnected"
- WorkerStatusUpdated = "WorkerStatusUpdated"
- WorkerDisconnected = "WorkerDisconnected"
-
- # # Timer Events
- # TimerCreated = "TimerCreated"
- # TimerFired = "TimerFired"
-
-
-class _BaseEvent[T: _EventType](BaseModel):
- """
- This is the event base-class, to please the Pydantic gods.
- PLEASE don't use this for anything unless you know why you are doing so,
- instead just use the events union :)
- """
-
- event_type: T
- event_id: EventId = EventId()
-
- def check_event_was_sent_by_correct_node(self, origin_id: NodeId) -> bool:
- """Check if the event was sent by the correct node.
-
- This is a placeholder implementation that always returns True.
- Subclasses can override this method to implement specific validation logic.
- """
- return True
diff --git a/shared/types/events/_events.py b/shared/types/events/_events.py
index 4f14b924..679bd940 100644
--- a/shared/types/events/_events.py
+++ b/shared/types/events/_events.py
@@ -1,20 +1,101 @@
-from typing import Annotated, Literal, Union
+import types
+from enum import Enum
+from typing import (
+ TYPE_CHECKING,
+ Annotated,
+ Literal,
+ Union,
+ get_args,
+ get_origin,
+ get_type_hints,
+)
from pydantic import Field
-from shared.topology import Connection, ConnectionProfile, Node, NodePerformanceProfile
+from shared.constants import get_error_reporting_message
+from shared.topology import Connection, ConnectionProfile, NodePerformanceProfile
from shared.types.common import NodeId
-from shared.types.events.chunks import GenerationChunk
+from shared.types.events.chunks import CommandId, GenerationChunk
from shared.types.tasks import Task, TaskId, TaskStatus
from shared.types.worker.common import InstanceId, NodeStatus
from shared.types.worker.instances import InstanceParams, TypeOfInstance
from shared.types.worker.runners import RunnerId, RunnerStatus
-from ._common import (
- CommandId,
- _BaseEvent, # pyright: ignore[reportPrivateUsage]
- _EventType, # pyright: ignore[reportPrivateUsage]
-)
+if TYPE_CHECKING:
+ pass
+
+from pydantic import BaseModel
+
+from shared.types.common import NewUUID
+
+
+class EventId(NewUUID):
+ """
+ Newtype around `NewUUID`
+ """
+
+
+# Event base-class boilerplate (you should basically never touch these)
+# Only very specialised registry or serialisation/deserialization logic might need know about these
+
+class _EventType(str, Enum):
+ """
+ Here are all the unique kinds of events that can be sent over the network.
+ """
+
+ # Task Saga Events
+ MLXInferenceSagaPrepare = "MLXInferenceSagaPrepare"
+ MLXInferenceSagaStartPrepare = "MLXInferenceSagaStartPrepare"
+
+ # Task Events
+ TaskCreated = "TaskCreated"
+ TaskStateUpdated = "TaskStateUpdated"
+ TaskDeleted = "TaskDeleted"
+
+ # Streaming Events
+ ChunkGenerated = "ChunkGenerated"
+
+ # Instance Events
+ InstanceCreated = "InstanceCreated"
+ InstanceDeleted = "InstanceDeleted"
+ InstanceActivated = "InstanceActivated"
+ InstanceDeactivated = "InstanceDeactivated"
+ InstanceReplacedAtomically = "InstanceReplacedAtomically"
+
+ # Runner Status Events
+ RunnerStatusUpdated = "RunnerStatusUpdated"
+
+ # Node Performance Events
+ NodePerformanceMeasured = "NodePerformanceMeasured"
+
+ # Topology Events
+ TopologyEdgeCreated = "TopologyEdgeCreated"
+ TopologyEdgeReplacedAtomically = "TopologyEdgeReplacedAtomically"
+ TopologyEdgeDeleted = "TopologyEdgeDeleted"
+ WorkerStatusUpdated = "WorkerStatusUpdated"
+
+ # # Timer Events
+ # TimerCreated = "TimerCreated"
+ # TimerFired = "TimerFired"
+
+
+class _BaseEvent[T: _EventType](BaseModel):
+ """
+ This is the event base-class, to please the Pydantic gods.
+ PLEASE don't use this for anything unless you know why you are doing so,
+ instead just use the events union :)
+ """
+
+ event_type: T
+ event_id: EventId = EventId()
+
+ def check_event_was_sent_by_correct_node(self, origin_id: NodeId) -> bool:
+ """Check if the event was sent by the correct node.
+
+ This is a placeholder implementation that always returns True.
+ Subclasses can override this method to implement specific validation logic.
+ """
+ return True
class TaskCreated(_BaseEvent[_EventType.TaskCreated]):
@@ -88,22 +169,12 @@ class NodePerformanceMeasured(_BaseEvent[_EventType.NodePerformanceMeasured]):
node_profile: NodePerformanceProfile
-class WorkerConnected(_BaseEvent[_EventType.WorkerConnected]):
- event_type: Literal[_EventType.WorkerConnected] = _EventType.WorkerConnected
- edge: Connection
-
-
class WorkerStatusUpdated(_BaseEvent[_EventType.WorkerStatusUpdated]):
event_type: Literal[_EventType.WorkerStatusUpdated] = _EventType.WorkerStatusUpdated
node_id: NodeId
node_state: NodeStatus
-class WorkerDisconnected(_BaseEvent[_EventType.WorkerDisconnected]):
- event_type: Literal[_EventType.WorkerDisconnected] = _EventType.WorkerDisconnected
- vertex_id: NodeId
-
-
class ChunkGenerated(_BaseEvent[_EventType.ChunkGenerated]):
event_type: Literal[_EventType.ChunkGenerated] = _EventType.ChunkGenerated
command_id: CommandId
@@ -112,7 +183,7 @@ class ChunkGenerated(_BaseEvent[_EventType.ChunkGenerated]):
class TopologyEdgeCreated(_BaseEvent[_EventType.TopologyEdgeCreated]):
event_type: Literal[_EventType.TopologyEdgeCreated] = _EventType.TopologyEdgeCreated
- vertex: Node
+ edge: Connection
class TopologyEdgeReplacedAtomically(_BaseEvent[_EventType.TopologyEdgeReplacedAtomically]):
@@ -136,9 +207,7 @@ _Event = Union[
InstanceReplacedAtomically,
RunnerStatusUpdated,
NodePerformanceMeasured,
- WorkerConnected,
WorkerStatusUpdated,
- WorkerDisconnected,
ChunkGenerated,
TopologyEdgeCreated,
TopologyEdgeReplacedAtomically,
@@ -151,6 +220,61 @@ Un-annotated union of all events. Only used internally to create the registry.
For all other usecases, use the annotated union of events :class:`Event` :)
"""
+
+def _check_event_type_consistency():
+ # Grab enum values from members
+ member_enum_values = [m for m in _EventType]
+
+ # grab enum values from the union => scrape the type annotation
+ union_enum_values: list[_EventType] = []
+ union_classes = list(get_args(_Event))
+ for cls in union_classes: # pyright: ignore[reportAny]
+ assert issubclass(cls, object), (
+ f"{get_error_reporting_message()}",
+ f"The class {cls} is NOT a subclass of {object}."
+ )
+
+ # ensure the first base parameter is ALWAYS _BaseEvent
+ base_cls = list(types.get_original_bases(cls))
+ assert len(base_cls) >= 1 and issubclass(base_cls[0], object) \
+ and issubclass(base_cls[0], _BaseEvent), (
+ f"{get_error_reporting_message()}",
+ f"The class {cls} does NOT inherit from {_BaseEvent} {get_origin(base_cls[0])}."
+ )
+
+ # grab type hints and extract the right values from it
+ cls_hints = get_type_hints(cls)
+ assert "event_type" in cls_hints and \
+ get_origin(cls_hints["event_type"]) is Literal, ( # pyright: ignore[reportAny]
+ f"{get_error_reporting_message()}",
+ f"The class {cls} is missing a {Literal}-annotated `event_type` field."
+ )
+
+ # make sure the value is an instance of `_EventType`
+ enum_value = list(get_args(cls_hints["event_type"]))
+ assert len(enum_value) == 1 and isinstance(enum_value[0], _EventType), (
+ f"{get_error_reporting_message()}",
+ f"The `event_type` of {cls} has a non-{_EventType} literal-type."
+ )
+ union_enum_values.append(enum_value[0])
+
+ # ensure there is a 1:1 bijection between the two
+ for m in member_enum_values:
+ assert m in union_enum_values, (
+ f"{get_error_reporting_message()}",
+ f"There is no event-type registered for {m} in {_Event}."
+ )
+ union_enum_values.remove(m)
+ assert len(union_enum_values) == 0, (
+ f"{get_error_reporting_message()}",
+ f"The following events have multiple event types defined in {_Event}: {union_enum_values}."
+ )
+
+
+_check_event_type_consistency()
+
+
+
Event = Annotated[_Event, Field(discriminator="event_type")]
"""Type of events, a discriminated union."""
diff --git a/shared/types/events/chunks.py b/shared/types/events/chunks.py
index e2cb7a7b..de5b079a 100644
--- a/shared/types/events/chunks.py
+++ b/shared/types/events/chunks.py
@@ -4,10 +4,15 @@ from typing import Annotated, Literal
from pydantic import BaseModel, Field, TypeAdapter
from shared.openai_compat import FinishReason
-from shared.types.events._common import CommandId
+from shared.types.common import NewUUID
from shared.types.models import ModelId
+class CommandId(NewUUID):
+ """
+ Newtype around `NewUUID` for command IDs
+ """
+
class ChunkType(str, Enum):
token = "token"
image = "image"
diff --git a/shared/types/events/commands.py b/shared/types/events/commands.py
index a4ec0e58..ae96f6d2 100644
--- a/shared/types/events/commands.py
+++ b/shared/types/events/commands.py
@@ -4,7 +4,8 @@ from typing import Annotated, Callable, Literal, Sequence
from pydantic import BaseModel, Field, TypeAdapter
from shared.types.api import ChatCompletionTaskParams
-from shared.types.events import CommandId, Event
+from shared.types.events import Event
+from shared.types.events.chunks import CommandId
from shared.types.state import InstanceId, State
diff --git a/shared/types/events/components.py b/shared/types/events/components.py
index f507d322..ddf9e30a 100644
--- a/shared/types/events/components.py
+++ b/shared/types/events/components.py
@@ -13,9 +13,10 @@ from typing import Callable
from pydantic import BaseModel, Field, model_validator
from shared.types.common import NodeId
-from shared.types.events._events import Event
from shared.types.state import State
+from ._events import Event
+
class EventFromEventLog[T: Event](BaseModel):
event: T
diff --git a/shared/types/topology.py b/shared/types/topology.py
index ce1d97ce..c41907ec 100644
--- a/shared/types/topology.py
+++ b/shared/types/topology.py
@@ -2,13 +2,10 @@ from typing import Iterable, Protocol
from pydantic import BaseModel, ConfigDict
-from shared.types.common import NewUUID, NodeId
+from shared.types.common import NodeId
from shared.types.profiling import ConnectionProfile, NodePerformanceProfile
-class ConnectionId(NewUUID):
- pass
-
class Connection(BaseModel):
source_node_id: NodeId
sink_node_id: NodeId
← 5097493a Fix tests
·
back to Exo
·
Fix the node-ID test 37301604 →