← back to Exo
refactor: Refactor most things
40793f1d8635a06c78e1ef28cb3c544ddeb1bf42 · 2025-07-02 21:11:49 +0100 · Arbion Halili
Files touched
M shared/pyproject.tomlM shared/types/events/common.pyM shared/types/events/events.pyA shared/types/graphs/common.pyM shared/types/graphs/resource_graph.pyM shared/types/models/sources.pyM shared/types/networking/edges.pyM shared/types/networking/services.pyM shared/types/networking/topology.pyM shared/types/states/master.pyM shared/types/states/shared.pyM shared/types/states/worker.pyM shared/types/tasks/common.pyM shared/types/worker/common.pyM shared/types/worker/instances.pyM shared/types/worker/runners.pyM shared/types/worker/shards.pyM uv.lock
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 →