[object Object]

← back to Exo

refactor: Refactor most things

40793f1d8635a06c78e1ef28cb3c544ddeb1bf42 · 2025-07-02 21:11:49 +0100 · Arbion Halili

Files touched

Diff

commit 40793f1d8635a06c78e1ef28cb3c544ddeb1bf42
Author: Arbion Halili <99731180+ToxicPine@users.noreply.github.com>
Date:   Wed Jul 2 21:11:49 2025 +0100

    refactor: Refactor most things
---
 shared/pyproject.toml                 |   1 +
 shared/types/events/common.py         |  38 +++++++++--
 shared/types/events/events.py         |  38 +++++++----
 shared/types/graphs/common.py         | 118 ++++++++++++++++++++++++++++++++++
 shared/types/graphs/resource_graph.py |   4 +-
 shared/types/models/sources.py        |  21 ++++--
 shared/types/networking/edges.py      |  89 +++++++++----------------
 shared/types/networking/services.py   |  25 ++++---
 shared/types/networking/topology.py   |  46 ++++++++-----
 shared/types/states/master.py         |  56 ++++++++++++++--
 shared/types/states/shared.py         |  16 ++++-
 shared/types/states/worker.py         |  20 +++---
 shared/types/tasks/common.py          |   6 +-
 shared/types/worker/common.py         |   2 +-
 shared/types/worker/instances.py      |  14 +++-
 shared/types/worker/runners.py        |   4 +-
 shared/types/worker/shards.py         |   4 --
 uv.lock                               |  11 ++++
 18 files changed, 373 insertions(+), 140 deletions(-)

diff --git a/shared/pyproject.toml b/shared/pyproject.toml
index d4ee919e..5721f6ad 100644
--- a/shared/pyproject.toml
+++ b/shared/pyproject.toml
@@ -10,6 +10,7 @@ dependencies = [
     "protobuf>=6.31.1",
     "pydantic>=2.11.7",
     "rich>=14.0.0",
+    "structlog>=25.4.0",
 ]
 
 [build-system]
diff --git a/shared/types/events/common.py b/shared/types/events/common.py
index f3d995a9..82fa3bc8 100644
--- a/shared/types/events/common.py
+++ b/shared/types/events/common.py
@@ -1,3 +1,4 @@
+import time
 from enum import Enum
 from typing import (
     Annotated,
@@ -11,9 +12,9 @@ from typing import (
     get_args,
 )
 
-from pydantic import BaseModel, Field, TypeAdapter
+from pydantic import BaseModel, Field, TypeAdapter, model_validator
 
-from shared.types.common import NewUUID
+from shared.types.common import NewUUID, NodeId
 
 
 class EventId(NewUUID):
@@ -39,11 +40,18 @@ class InstanceEventTypes(str, Enum):
     InstanceCreated = "InstanceCreated"
     InstanceDeleted = "InstanceDeleted"
     InstanceReplacedAtomically = "InstanceReplacedAtomically"
+    InstanceStatusUpdated = "InstanceStatusUpdated"
+
+
+class InstanceStateEventTypes(str, Enum):
     InstanceRunnerStateUpdated = "InstanceRunnerStateUpdated"
 
 
-class NodeEventTypes(str, Enum):
-    NodeStateUpdated = "NodeStateUpdated"
+class NodeStatusEventTypes(str, Enum):
+    NodeStatusUpdated = "NodeStatusUpdated"
+
+
+class NodeProfileEventTypes(str, Enum):
     NodeProfileUpdated = "NodeProfileUpdated"
 
 
@@ -62,7 +70,9 @@ EventTypes = Union[
     TaskEventTypes,
     StreamingEventTypes,
     InstanceEventTypes,
-    NodeEventTypes,
+    InstanceStateEventTypes,
+    NodeStatusEventTypes,
+    NodeProfileEventTypes,
     EdgeEventTypes,
     TimerEventTypes,
     MLXEventTypes,
@@ -72,11 +82,27 @@ EventTypeT = TypeVar("EventTypeT", bound=EventTypes)
 TEventType = TypeVar("TEventType", bound=EventTypes, covariant=True)
 
 
-class Event(BaseModel, Generic[TEventType]):
+class SecureEventProtocol(Protocol):
+    def check_origin_id(self, origin_id: NodeId) -> bool: ...
+
+
+class Event(BaseModel, SecureEventProtocol, Generic[TEventType]):
     event_type: TEventType
     event_id: EventId
 
 
+class WrappedEvent(BaseModel, Generic[TEventType]):
+    event: Event[TEventType]
+    origin_id: NodeId
+    origin_timestamp: int = Field(default_factory=lambda: int(time.time()))
+
+    @model_validator(mode="after")
+    def check_origin_id(self) -> "WrappedEvent[TEventType]":
+        if self.event.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]
     sequence_number: int = Field(gt=0)
diff --git a/shared/types/events/events.py b/shared/types/events/events.py
index db5a3e32..22a6dd89 100644
--- a/shared/types/events/events.py
+++ b/shared/types/events/events.py
@@ -8,8 +8,10 @@ from shared.types.common import NewUUID, NodeId
 from shared.types.events.common import (
     Event,
     InstanceEventTypes,
+    InstanceStateEventTypes,
     MLXEventTypes,
-    NodeEventTypes,
+    NodeProfileEventTypes,
+    NodeStatusEventTypes,
     StreamingEventTypes,
     TaskEventTypes,
     TimerEventTypes,
@@ -22,8 +24,8 @@ from shared.types.tasks.common import (
     TaskType,
     TaskUpdate,
 )
-from shared.types.worker.common import InstanceId, NodeState
-from shared.types.worker.instances import InstanceData
+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
 
 
@@ -73,9 +75,19 @@ class InstanceDeleted(Event[InstanceEventTypes.InstanceDeleted]):
     instance_id: InstanceId
 
 
-class InstanceRunnerStateUpdated(Event[InstanceEventTypes.InstanceRunnerStateUpdated]):
-    event_type: Literal[InstanceEventTypes.InstanceRunnerStateUpdated] = (
-        InstanceEventTypes.InstanceRunnerStateUpdated
+class InstanceStatusUpdated(Event[InstanceEventTypes.InstanceStatusUpdated]):
+    event_type: Literal[InstanceEventTypes.InstanceStatusUpdated] = (
+        InstanceEventTypes.InstanceStatusUpdated
+    )
+    instance_id: InstanceId
+    instance_status: InstanceStatus
+
+
+class InstanceRunnerStateUpdated(
+    Event[InstanceStateEventTypes.InstanceRunnerStateUpdated]
+):
+    event_type: Literal[InstanceStateEventTypes.InstanceRunnerStateUpdated] = (
+        InstanceStateEventTypes.InstanceRunnerStateUpdated
     )
     instance_id: InstanceId
     state_update: Tuple[RunnerId, RunnerState[RunnerStateType]]
@@ -106,20 +118,20 @@ class MLXInferenceSagaStartPrepare(Event[MLXEventTypes.MLXInferenceSagaStartPrep
     instance_id: InstanceId
 
 
-class NodeProfileUpdated(Event[NodeEventTypes.NodeProfileUpdated]):
-    event_type: Literal[NodeEventTypes.NodeProfileUpdated] = (
-        NodeEventTypes.NodeProfileUpdated
+class NodeProfileUpdated(Event[NodeProfileEventTypes.NodeProfileUpdated]):
+    event_type: Literal[NodeProfileEventTypes.NodeProfileUpdated] = (
+        NodeProfileEventTypes.NodeProfileUpdated
     )
     node_id: NodeId
     node_profile: NodeProfile
 
 
-class NodeStateUpdated(Event[NodeEventTypes.NodeStateUpdated]):
-    event_type: Literal[NodeEventTypes.NodeStateUpdated] = (
-        NodeEventTypes.NodeStateUpdated
+class NodeStatusUpdated(Event[NodeStatusEventTypes.NodeStatusUpdated]):
+    event_type: Literal[NodeStatusEventTypes.NodeStatusUpdated] = (
+        NodeStatusEventTypes.NodeStatusUpdated
     )
     node_id: NodeId
-    node_state: NodeState
+    node_state: NodeStatus
 
 
 class ChunkGenerated(Event[StreamingEventTypes.ChunkGenerated]):
diff --git a/shared/types/graphs/common.py b/shared/types/graphs/common.py
new file mode 100644
index 00000000..878d6d35
--- /dev/null
+++ b/shared/types/graphs/common.py
@@ -0,0 +1,118 @@
+from collections.abc import Mapping
+from typing import Generic, Protocol, Set, Tuple, TypeVar, overload
+
+from pydantic import BaseModel
+
+from shared.types.common import NewUUID
+
+EdgeTypeT = TypeVar("EdgeTypeT", covariant=True)
+VertexTypeT = TypeVar("VertexTypeT", covariant=True)
+EdgeIdT = TypeVar("EdgeIdT", bound=NewUUID)
+VertexIdT = TypeVar("VertexIdT", bound=NewUUID)
+
+
+class VertexData(BaseModel, Generic[VertexTypeT]):
+    vertex_type: VertexTypeT
+
+
+class EdgeData(BaseModel, Generic[EdgeTypeT]):
+    edge_type: EdgeTypeT
+
+
+class BaseEdge(BaseModel, Generic[EdgeTypeT, EdgeIdT, VertexIdT]):
+    edge_vertices: Tuple[VertexIdT, VertexIdT]
+    edge_data: EdgeData[EdgeTypeT]
+
+
+class BaseVertex(BaseModel, Generic[VertexTypeT, EdgeIdT]):
+    vertex_data: VertexData[VertexTypeT]
+
+
+class Vertex(
+    BaseVertex[VertexTypeT, EdgeIdT], Generic[VertexTypeT, EdgeIdT, VertexIdT]
+):
+    vertex_id: VertexIdT
+
+
+class Edge(
+    BaseEdge[EdgeTypeT, EdgeIdT, VertexIdT], Generic[EdgeTypeT, EdgeIdT, VertexIdT]
+):
+    edge_id: EdgeIdT
+
+
+class GraphData(BaseModel, Generic[EdgeTypeT, VertexTypeT, EdgeIdT, VertexIdT]):
+    edges: Mapping[EdgeIdT, EdgeData[EdgeTypeT]]
+    vertices: Mapping[VertexIdT, VertexData[VertexTypeT]]
+
+
+class GraphProtocol(Protocol, Generic[EdgeTypeT, VertexTypeT, EdgeIdT, VertexIdT]):
+    def list_edges(self) -> Set[EdgeIdT]: ...
+    def list_vertices(self) -> Set[VertexIdT]: ...
+    def get_vertices_from_edges(
+        self, edges: Set[EdgeIdT]
+    ) -> Mapping[EdgeIdT, Set[VertexIdT]]: ...
+    def get_edges_from_vertices(
+        self, vertices: Set[VertexIdT]
+    ) -> Mapping[VertexIdT, Set[EdgeIdT]]: ...
+    def get_edge_data(
+        self, edges: Set[EdgeIdT]
+    ) -> Mapping[EdgeIdT, EdgeData[EdgeTypeT]]: ...
+    def get_vertex_data(
+        self, vertices: Set[VertexIdT]
+    ) -> Mapping[VertexIdT, VertexData[VertexTypeT]]: ...
+
+
+class UpdatableGraphProtocol(GraphProtocol[EdgeTypeT, VertexTypeT, EdgeIdT, VertexIdT]):
+    def check_edges_exists(self, edge_id: EdgeIdT) -> bool: ...
+    def check_vertex_exists(self, vertex_id: VertexIdT) -> bool: ...
+    def _add_edge(self, edge_id: EdgeIdT, edge_data: EdgeData[EdgeTypeT]) -> None: ...
+    def _add_vertex(
+        self, vertex_id: VertexIdT, vertex_data: VertexData[VertexTypeT]
+    ) -> None: ...
+    def _remove_edge(self, edge_id: EdgeIdT) -> None: ...
+    def _remove_vertex(self, vertex_id: VertexIdT) -> None: ...
+    ###
+    @overload
+    def attach_edge(self, edge: Edge[EdgeTypeT, EdgeIdT, VertexIdT]) -> None: ...
+    @overload
+    def attach_edge(
+        self,
+        edge: Edge[EdgeTypeT, EdgeIdT, VertexIdT],
+        extra_vertex: Vertex[VertexTypeT, EdgeIdT, VertexIdT],
+    ) -> None: ...
+    def attach_edge(
+        self,
+        edge: Edge[EdgeTypeT, EdgeIdT, VertexIdT],
+        extra_vertex: Vertex[VertexTypeT, EdgeIdT, VertexIdT] | None = None,
+    ) -> None:
+        base_vertex = edge.edge_vertices[0]
+        target_vertex = edge.edge_vertices[1]
+        base_vertex_exists = self.check_vertex_exists(base_vertex)
+        target_vertex_exists = self.check_vertex_exists(target_vertex)
+
+        if not base_vertex_exists:
+            raise ValueError("Base Vertex Does Not Exist")
+
+        match (target_vertex_exists, extra_vertex is not None):
+            case (True, False):
+                raise ValueError("New Vertex Already Exists")
+            case (False, True):
+                if extra_vertex is None:
+                    raise ValueError("BUG: Extra Vertex Must Be Provided")
+                self._add_vertex(extra_vertex.vertex_id, extra_vertex.vertex_data)
+            case (False, False):
+                raise ValueError(
+                    "New Vertex Must Be Provided For Non-Existent Target Vertex"
+                )
+            case (True, True):
+                raise ValueError("New Vertex Already Exists")
+
+        self._add_edge(edge.edge_id, edge.edge_data)
+
+
+class Graph(
+    BaseModel,
+    Generic[EdgeTypeT, VertexTypeT, EdgeIdT, VertexIdT],
+    GraphProtocol[EdgeTypeT, VertexTypeT, EdgeIdT, VertexIdT],
+):
+    graph_data: GraphData[EdgeTypeT, VertexTypeT, EdgeIdT, VertexIdT]
diff --git a/shared/types/graphs/resource_graph.py b/shared/types/graphs/resource_graph.py
index 25f7dd52..4c469d9c 100644
--- a/shared/types/graphs/resource_graph.py
+++ b/shared/types/graphs/resource_graph.py
@@ -5,7 +5,7 @@ from pydantic import BaseModel
 from shared.types.common import NodeId
 from shared.types.networking.topology import Topology
 from shared.types.profiling.common import NodeProfile
-from shared.types.worker.common import NodeState
+from shared.types.worker.common import NodeStatus
 
 
 class ResourceGraph(BaseModel): ...
@@ -13,6 +13,6 @@ class ResourceGraph(BaseModel): ...
 
 def get_graph_of_compute_resources(
     network_topology: Topology,
-    node_states: Mapping[NodeId, NodeState],
+    node_statuses: Mapping[NodeId, NodeStatus],
     node_profiles: Mapping[NodeId, NodeProfile],
 ) -> ResourceGraph: ...
diff --git a/shared/types/models/sources.py b/shared/types/models/sources.py
index 419ed264..8f636a26 100644
--- a/shared/types/models/sources.py
+++ b/shared/types/models/sources.py
@@ -11,14 +11,20 @@ class SourceType(str, Enum):
     GitHub = "GitHub"
 
 
+class SourceFormatType(str, Enum):
+    HuggingFaceTransformers = "HuggingFaceTransformers"
+
+
 T = TypeVar("T", bound=SourceType)
+S = TypeVar("S", bound=SourceFormatType)
 
 RepoPath = Annotated[str, Field(pattern=r"^[^/]+/[^/]+$")]
 
 
-class BaseModelSource(BaseModel, Generic[T]):
+class BaseModelSource(BaseModel, Generic[T, S]):
     model_uuid: ModelId
     source_type: T
+    source_format: S
     source_data: Any
 
 
@@ -33,13 +39,18 @@ class GitHubModelSourceData(BaseModel):
 
 
 @final
-class HuggingFaceModelSource(BaseModelSource[SourceType.HuggingFace]):
+class HuggingFaceModelSource(
+    BaseModelSource[SourceType.HuggingFace, SourceFormatType.HuggingFaceTransformers]
+):
     source_type: Literal[SourceType.HuggingFace] = SourceType.HuggingFace
+    source_format: Literal[SourceFormatType.HuggingFaceTransformers] = (
+        SourceFormatType.HuggingFaceTransformers
+    )
     source_data: HuggingFaceModelSourceData
 
 
 @final
-class GitHubModelSource(BaseModelSource[SourceType.GitHub]):
+class GitHubModelSource(BaseModelSource[SourceType.GitHub, S]):
     source_type: Literal[SourceType.GitHub] = SourceType.GitHub
     source_data: GitHubModelSourceData
 
@@ -47,9 +58,9 @@ class GitHubModelSource(BaseModelSource[SourceType.GitHub]):
 _ModelSource = Annotated[
     Union[
         HuggingFaceModelSource,
-        GitHubModelSource,
+        GitHubModelSource[SourceFormatType.HuggingFaceTransformers],
     ],
     Field(discriminator="source_type"),
 ]
-ModelSource = BaseModelSource[SourceType]
+ModelSource = BaseModelSource[SourceType, SourceFormatType]
 ModelSourceAdapter: TypeAdapter[ModelSource] = TypeAdapter(_ModelSource)
diff --git a/shared/types/networking/edges.py b/shared/types/networking/edges.py
index 0977caf1..3c90a837 100644
--- a/shared/types/networking/edges.py
+++ b/shared/types/networking/edges.py
@@ -1,10 +1,13 @@
-from collections.abc import Mapping
 from enum import Enum
-from typing import Annotated, Generic, NamedTuple, TypeVar, final
+from typing import Generic, Mapping, Tuple, TypeVar, final
 
-from pydantic import AfterValidator, BaseModel, IPvAnyAddress
+from pydantic import BaseModel, IPvAnyAddress
 
 from shared.types.common import NewUUID, NodeId
+from shared.types.graphs.common import (
+    Edge,
+    EdgeData,
+)
 
 
 class EdgeId(NewUUID):
@@ -30,78 +33,50 @@ class EdgeDataTransferRate(BaseModel):
     jitter: float
 
 
-class EdgeMetadata(BaseModel, Generic[AdP, ApP]): ...
+class NetworkEdgeMetadata(BaseModel, Generic[AdP, ApP]): ...
 
 
 @final
-class EdgeType(BaseModel, Generic[AdP, ApP]):
+class NetworkEdgeType(BaseModel, Generic[AdP, ApP]):
     addressing_protocol: AdP
     application_protocol: ApP
 
 
 @final
-class EdgeDirection(NamedTuple):
-    source: NodeId
-    sink: NodeId
-
-
-@final
-class MLXEdgeContext(EdgeMetadata[AddressingProtocol.IPvAny, ApplicationProtocol.MLX]):
+class MLXEdgeContext(
+    NetworkEdgeMetadata[AddressingProtocol.IPvAny, ApplicationProtocol.MLX]
+):
     source_ip: IPvAnyAddress
     sink_ip: IPvAnyAddress
 
 
-class EdgeDataType(str, Enum):
-    DISCOVERED = "discovered"
-    PROFILED = "profiled"
-    UNKNOWN = "unknown"
+class NetworkEdgeInfoType(str, Enum):
+    network_profile = "network_profile"
+    other = "other"
+
+
+AllNetworkEdgeInfo = Tuple[NetworkEdgeInfoType.network_profile]
 
 
-EdgeDataTypeT = TypeVar("EdgeDataTypeT", bound=EdgeDataType)
+NetworkEdgeInfoTypeT = TypeVar(
+    "NetworkEdgeInfoTypeT", bound=NetworkEdgeInfoType, covariant=True
+)
+
+
+class NetworkEdgeInfo(BaseModel, Generic[NetworkEdgeInfoTypeT]):
+    edge_info_type: NetworkEdgeInfoTypeT
 
 
-class EdgeData(BaseModel, Generic[EdgeDataTypeT]):
-    edge_data_type: EdgeDataTypeT
+SetOfEdgeInfo = TypeVar("SetOfEdgeInfo", bound=Tuple[NetworkEdgeInfoType, ...])
 
 
-class EdgeProfile(EdgeData[EdgeDataType.PROFILED]):
+class NetworkEdgeData(EdgeData[NetworkEdgeType[AdP, ApP]], Generic[AdP, ApP]):
+    edge_info: Mapping[NetworkEdgeInfoType, NetworkEdgeInfo[NetworkEdgeInfoType]]
+    edge_metadata: NetworkEdgeMetadata[AdP, ApP]
+
+
+class NetworkEdgeProfile(NetworkEdgeInfo[NetworkEdgeInfoTypeT]):
     edge_data_transfer_rate: EdgeDataTransferRate
 
 
-def validate_mapping(
-    edge_data: Mapping[EdgeDataType, EdgeData[EdgeDataType]],
-) -> Mapping[EdgeDataType, EdgeData[EdgeDataType]]:
-    """Validates that each EdgeData value has an edge_data_type matching its key."""
-    for key, value in edge_data.items():
-        if key != value.edge_data_type:
-            raise ValueError(
-                f"Edge Data Type Mismatch: key {key} != value {value.edge_data_type}"
-            )
-    return edge_data
-
-
-class Edge(BaseModel, Generic[AdP, ApP, EdgeDataTypeT]):
-    edge_type: EdgeType[AdP, ApP]
-    edge_direction: EdgeDirection
-    edge_data: Annotated[
-        Mapping[EdgeDataType, EdgeData[EdgeDataType]], AfterValidator(validate_mapping)
-    ]
-    edge_metadata: EdgeMetadata[AdP, ApP]
-
-
-"""
-an_edge: UniqueEdge[Literal[AddressingProtocol.IPvAny], Literal[ApplicationProtocol.MLX]] = UniqueEdge(
-    edge_identifier=EdgeId(UUID().hex),
-    edge_info=ProfiledEdge(
-        edge_direction=EdgeDirection(source=NodeId("1"), sink=NodeId("2")),
-        edge_type=EdgeType(
-            addressing_protocol=AddressingProtocol.IPvAny,
-            application_protocol=ApplicationProtocol.MLX
-        ),
-        edge_data=EdgeData(
-            edge_data_transfer_rate=EdgeDataTransferRate(throughput=1000, latency=0.1, jitter=0.01)
-        ),
-        edge_metadata=MLXEdgeContext(source_ip=IPv4Address("192.168.1.1"), sink_ip=IPv4Address("192.168.1.2"))
-    )
-)
-"""
+class NetworkEdge(Edge[NetworkEdgeType[AdP, ApP], EdgeId, NodeId]): ...
diff --git a/shared/types/networking/services.py b/shared/types/networking/services.py
index 119defc9..f7319c43 100644
--- a/shared/types/networking/services.py
+++ b/shared/types/networking/services.py
@@ -1,27 +1,26 @@
-from typing import Annotated, Callable, NewType, Protocol
+from typing import Callable, NewType, Protocol, TypeVar
 
-from pydantic import BaseModel, Field
-
-from shared.types.common import NodeId
 from shared.types.networking.edges import (
     AddressingProtocol,
     ApplicationProtocol,
-    Edge,
-    EdgeDataType,
     EdgeId,
+    NetworkEdge,
 )
 
 TopicName = NewType("TopicName", str)
 
-
-class WrappedMessage(BaseModel):
-    node_id: NodeId
-    unix_timestamp: Annotated[int, Field(gt=0)]
+MessageT = TypeVar("MessageT", bound=object)
 
 
-PubSubMessageHandler = Callable[[TopicName, WrappedMessage], None]
+PubSubMessageHandler = Callable[[TopicName, MessageT], None]
 NodeConnectedHandler = Callable[
-    [EdgeId, Edge[AddressingProtocol, ApplicationProtocol, EdgeDataType.DISCOVERED]],
+    [
+        EdgeId,
+        NetworkEdge[
+            AddressingProtocol,
+            ApplicationProtocol,
+        ],
+    ],
     None,
 ]
 NodeDisconnectedHandler = Callable[[EdgeId], None]
@@ -38,6 +37,6 @@ class DiscoveryService(Protocol):
 
 class PubSubService(Protocol):
     def register_handler(
-        self, key: str, topic_name: TopicName, handler: PubSubMessageHandler
+        self, key: str, topic_name: TopicName, handler: PubSubMessageHandler[MessageT]
     ) -> None: ...
     def deregister_handler(self, key: str) -> None: ...
diff --git a/shared/types/networking/topology.py b/shared/types/networking/topology.py
index 1f0c8144..6768c15c 100644
--- a/shared/types/networking/topology.py
+++ b/shared/types/networking/topology.py
@@ -1,28 +1,40 @@
-from collections.abc import Mapping, Sequence
-from typing import Literal
-
-from pydantic import BaseModel
-
+from shared.types.common import NodeId
+from shared.types.graphs.common import Graph, GraphData
 from shared.types.networking.edges import (
     AddressingProtocol,
     ApplicationProtocol,
-    Edge,
-    EdgeDataType,
     EdgeId,
+    NetworkEdge,
 )
 
 
-class Topology(BaseModel):
-    edges: Mapping[
+class Topology(
+    Graph[
+        NetworkEdge[AddressingProtocol, ApplicationProtocol],
+        None,
         EdgeId,
-        Edge[AddressingProtocol, ApplicationProtocol, Literal[EdgeDataType.DISCOVERED]],
+        NodeId,
+    ]
+):
+    graph_data: GraphData[
+        NetworkEdge[AddressingProtocol, ApplicationProtocol],
+        None,
+        EdgeId,
+        NodeId,
     ]
 
 
-class EdgeMap(BaseModel):
-    edges: Mapping[EdgeId, Edge[AddressingProtocol, ApplicationProtocol, EdgeDataType]]
-
-
-class NetworkState(BaseModel):
-    topology: Topology
-    history: Sequence[Topology]
+class OrphanedPartOfTopology(
+    Graph[
+        NetworkEdge[AddressingProtocol, ApplicationProtocol],
+        None,
+        EdgeId,
+        NodeId,
+    ]
+):
+    graph_data: GraphData[
+        NetworkEdge[AddressingProtocol, ApplicationProtocol],
+        None,
+        EdgeId,
+        NodeId,
+    ]
diff --git a/shared/types/states/master.py b/shared/types/states/master.py
index ca11ae32..5f47ec18 100644
--- a/shared/types/states/master.py
+++ b/shared/types/states/master.py
@@ -1,27 +1,72 @@
 from collections.abc import Mapping, Sequence
+from enum import Enum
 from queue import Queue
+from typing import Generic, TypeVar
 
 from pydantic import BaseModel
 
 from shared.types.common import NodeId
-from shared.types.events.common import Event, EventTypes
+from shared.types.events.common import (
+    EdgeEventTypes,
+    Event,
+    EventTypes,
+    NodeProfileEventTypes,
+    NodeStatusEventTypes,
+    State,
+)
 from shared.types.graphs.resource_graph import ResourceGraph
-from shared.types.networking.topology import NetworkState
+from shared.types.networking.edges import (
+    AddressingProtocol,
+    ApplicationProtocol,
+    EdgeId,
+    NetworkEdge,
+)
+from shared.types.networking.topology import OrphanedPartOfTopology, Topology
 from shared.types.profiling.common import NodeProfile
 from shared.types.states.shared import SharedState
-from shared.types.worker.common import NodeState
+from shared.types.worker.common import NodeStatus
 from shared.types.worker.instances import InstanceData, InstanceId
 
 
 class ExternalCommand(BaseModel): ...
 
 
+class CachePolicyType(str, Enum):
+    KeepAll = "KeepAll"
+
+
+CachePolicyTypeT = TypeVar("CachePolicyTypeT", bound=CachePolicyType)
+
+
+class CachePolicy(BaseModel, Generic[CachePolicyTypeT]):
+    policy_type: CachePolicyTypeT
+
+
+class NodeProfileState(State[NodeProfileEventTypes]):
+    node_profiles: Mapping[NodeId, NodeProfile]
+
+
+class NodeStatusState(State[NodeStatusEventTypes]):
+    node_status: Mapping[NodeId, NodeStatus]
+
+
+class NetworkState(State[EdgeEventTypes]):
+    topology: Topology
+    history: Sequence[OrphanedPartOfTopology]
+
+    def delete_edge(self, edge_id: EdgeId) -> None: ...
+    def add_edge(
+        self, edge: NetworkEdge[AddressingProtocol, ApplicationProtocol]
+    ) -> None: ...
+
+
 class MasterState(SharedState):
     network_state: NetworkState
-    node_profiles: Mapping[NodeId, NodeProfile]
-    node_states: Mapping[NodeId, NodeState]
+    node_profiles: NodeProfileState
+    node_status: NodeStatusState
     job_inbox: Queue[ExternalCommand]
     job_outbox: Queue[ExternalCommand]
+    cache_policy: CachePolicy[CachePolicyType]
 
 
 def get_inference_plan(
@@ -29,6 +74,7 @@ def get_inference_plan(
     outbox: Queue[ExternalCommand],
     resource_graph: ResourceGraph,
     current_instances: Mapping[InstanceId, InstanceData],
+    cache_policy: CachePolicy[CachePolicyType],
 ) -> Mapping[InstanceId, InstanceData]: ...
 
 
diff --git a/shared/types/states/shared.py b/shared/types/states/shared.py
index acf09499..1dae6823 100644
--- a/shared/types/states/shared.py
+++ b/shared/types/states/shared.py
@@ -1,14 +1,24 @@
 from collections.abc import Mapping
+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.tasks.common import Task, TaskId, TaskType
 from shared.types.worker.common import InstanceId
-from shared.types.worker.instances import InstanceData
+from shared.types.worker.instances import BaseInstance
+
+
+class Instances(State[InstanceStateEventTypes]):
+    instances: Mapping[InstanceId, BaseInstance]
+
+
+class Tasks(State[TaskEventTypes]):
+    tasks: Mapping[TaskId, Task[TaskType]]
 
 
 class SharedState(BaseModel):
     node_id: NodeId
-    compute_instances: Mapping[InstanceId, InstanceData]
-    compute_tasks: dict[TaskId, Task[TaskType]]
+    known_instances: Instances
+    compute_tasks: Sequence[Task[TaskType]]
diff --git a/shared/types/states/worker.py b/shared/types/states/worker.py
index 37a187da..5db788df 100644
--- a/shared/types/states/worker.py
+++ b/shared/types/states/worker.py
@@ -1,15 +1,17 @@
 from collections.abc import Mapping
-from typing import Tuple
 
-from shared.types.models.common import ModelId
+from shared.types.common import NodeId
+from shared.types.events.common import (
+    NodeStatusEventTypes,
+    State,
+)
 from shared.types.states.shared import SharedState
-from shared.types.worker.common import NodeState
-from shared.types.worker.downloads import BaseDownloadProgress, DownloadStatus
-from shared.types.worker.shards import ShardData, ShardType
+from shared.types.worker.common import NodeStatus
+
+
+class NodeStatusState(State[NodeStatusEventTypes]):
+    node_status: Mapping[NodeId, NodeStatus]
 
 
 class WorkerState(SharedState):
-    node_state: NodeState
-    download_state: Mapping[
-        Tuple[ModelId, ShardData[ShardType]], BaseDownloadProgress[DownloadStatus]
-    ]
+    node_status: NodeStatusState
diff --git a/shared/types/tasks/common.py b/shared/types/tasks/common.py
index 4baf87fb..da99804f 100644
--- a/shared/types/tasks/common.py
+++ b/shared/types/tasks/common.py
@@ -74,7 +74,11 @@ class FailedTask(TaskUpdate[TaskStatusType.Failed]):
     error_message: Mapping[RunnerId, str]
 
 
-class Task(BaseModel):
+class BaseTask(BaseModel):
     task_data: TaskData[TaskType]
     task_status: TaskUpdate[TaskStatusType]
     on_instance: InstanceId
+
+
+class Task(BaseTask):
+    task_id: TaskId
diff --git a/shared/types/worker/common.py b/shared/types/worker/common.py
index 0d53ddc5..5fa78f74 100644
--- a/shared/types/worker/common.py
+++ b/shared/types/worker/common.py
@@ -11,7 +11,7 @@ class RunnerId(NewUUID):
     pass
 
 
-class NodeState(str, Enum):
+class NodeStatus(str, Enum):
     Idle = "Idle"
     Running = "Running"
     Paused = "Paused"
diff --git a/shared/types/worker/instances.py b/shared/types/worker/instances.py
index 0a3f8728..04884d14 100644
--- a/shared/types/worker/instances.py
+++ b/shared/types/worker/instances.py
@@ -1,4 +1,5 @@
 from collections.abc import Mapping
+from enum import Enum
 
 from pydantic import BaseModel
 
@@ -11,6 +12,11 @@ from shared.types.worker.runners import (
 )
 
 
+class InstanceStatus(str, Enum):
+    ACTIVE = "active"
+    INACTIVE = "inactive"
+
+
 class InstanceState(BaseModel):
     runner_states: Mapping[RunnerId, RunnerState[RunnerStateType]]
 
@@ -19,7 +25,11 @@ class InstanceData(BaseModel):
     runner_placements: RunnerPlacement
 
 
-class Instance(BaseModel):
-    instance_id: InstanceId
+class BaseInstance(BaseModel):
     instance_data: InstanceData
     instance_state: InstanceState
+    instance_status: InstanceStatus
+
+
+class Instance(BaseInstance):
+    instance_id: InstanceId
diff --git a/shared/types/worker/runners.py b/shared/types/worker/runners.py
index decf349f..1ca1dc22 100644
--- a/shared/types/worker/runners.py
+++ b/shared/types/worker/runners.py
@@ -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 Shard, ShardType
+from shared.types.worker.shards import ShardData, ShardType
 
 
 class RunnerStateType(str, Enum):
@@ -57,7 +57,7 @@ class RunnerData(BaseModel):
 
 class RunnerPlacement(BaseModel):
     model_id: ModelId
-    runner_to_shard: Mapping[RunnerId, Shard[ShardType]]
+    runner_to_shard: Mapping[RunnerId, ShardData[ShardType]]
     node_to_runner: Mapping[NodeId, Sequence[RunnerId]]
 
     @model_validator(mode="after")
diff --git a/shared/types/worker/shards.py b/shared/types/worker/shards.py
index 3e9055ae..f7a97a42 100644
--- a/shared/types/worker/shards.py
+++ b/shared/types/worker/shards.py
@@ -13,7 +13,3 @@ ShardTypeT = TypeVar("ShardTypeT", bound=ShardType)
 
 class ShardData(BaseModel, Generic[ShardTypeT]):
     shard_type: ShardTypeT
-
-
-class Shard(BaseModel, Generic[ShardTypeT]):
-    shard_data: ShardData[ShardTypeT]
diff --git a/uv.lock b/uv.lock
index d08efbb3..b76dd752 100644
--- a/uv.lock
+++ b/uv.lock
@@ -141,6 +141,7 @@ dependencies = [
     { name = "protobuf", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" },
     { name = "pydantic", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" },
     { name = "rich", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" },
+    { name = "structlog", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" },
 ]
 
 [package.dev-dependencies]
@@ -155,6 +156,7 @@ requires-dist = [
     { name = "protobuf", specifier = ">=6.31.1" },
     { name = "pydantic", specifier = ">=2.11.7" },
     { name = "rich", specifier = ">=14.0.0" },
+    { name = "structlog", specifier = ">=25.4.0" },
 ]
 
 [package.metadata.requires-dev]
@@ -486,6 +488,15 @@ wheels = [
     { url = "https://files.pythonhosted.org/packages/e9/44/75a9c9421471a6c4805dbf2356f7c181a29c1879239abab1ea2cc8f38b40/sniffio-1.3.1-py3-none-any.whl", hash = "sha256:2f6da418d1f1e0fddd844478f41680e794e6051915791a034ff65e5f100525a2", size = 10235, upload-time = "2024-02-25T23:20:01.196Z" },
 ]
 
+[[package]]
+name = "structlog"
+version = "25.4.0"
+source = { registry = "https://pypi.org/simple" }
+sdist = { url = "https://files.pythonhosted.org/packages/79/b9/6e672db4fec07349e7a8a8172c1a6ae235c58679ca29c3f86a61b5e59ff3/structlog-25.4.0.tar.gz", hash = "sha256:186cd1b0a8ae762e29417095664adf1d6a31702160a46dacb7796ea82f7409e4", size = 1369138, upload-time = "2025-06-02T08:21:12.971Z" }
+wheels = [
+    { url = "https://files.pythonhosted.org/packages/a0/4a/97ee6973e3a73c74c8120d59829c3861ea52210667ec3e7a16045c62b64d/structlog-25.4.0-py3-none-any.whl", hash = "sha256:fe809ff5c27e557d14e613f45ca441aabda051d119ee5a0102aaba6ce40eed2c", size = 68720, upload-time = "2025-06-02T08:21:11.43Z" },
+]
+
 [[package]]
 name = "tqdm"
 version = "4.67.1"

← 8596d5c5 refactor: Fix UUID implementation  ·  back to Exo  ·  feature: Simplest utilities for logging 7dd8a979 →