[object Object]

← 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

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 →