[object Object]

← back to Exo

Refactor worker + master state into single state

bae58dd368486b3d829ba57f1513ddcdaf6d6ca7 · 2025-07-21 19:36:54 +0100 · Seth Howes

Files touched

Diff

commit bae58dd368486b3d829ba57f1513ddcdaf6d6ca7
Author: Seth Howes <71157822+sethhowes@users.noreply.github.com>
Date:   Mon Jul 21 19:36:54 2025 +0100

    Refactor worker + master state into single state
---
 master/main.py                   | 12 ++++----
 master/placement.py              |  4 +--
 shared/types/models/model.py     |  1 -
 shared/types/state.py            | 41 ++++++++++++++++++++++++++++
 shared/types/states/master.py    | 59 ----------------------------------------
 shared/types/states/shared.py    | 52 -----------------------------------
 shared/types/states/worker.py    | 21 --------------
 worker/main.py                   |  8 +++---
 worker/test_worker_state.py      | 15 +++++-----
 worker/tests/conftest.py         | 18 ++++++------
 worker/tests/test_worker_plan.py | 20 ++++++--------
 11 files changed, 77 insertions(+), 174 deletions(-)

diff --git a/master/main.py b/master/main.py
index a81ccd91..730289ac 100644
--- a/master/main.py
+++ b/master/main.py
@@ -24,19 +24,19 @@ from shared.types.events.common import (
 )
 from shared.types.models.common import ModelId
 from shared.types.models.model import ModelInfo
-from shared.types.states.master import MasterState
+from shared.types.state import State
 from shared.types.worker.common import InstanceId
 from shared.types.worker.instances import Instance
 
 
 # Restore State
-def get_master_state(logger: Logger) -> MasterState:
+def get_state(logger: Logger) -> State:
     if EXO_MASTER_STATE.exists():
         with open(EXO_MASTER_STATE, "r") as f:
-            return MasterState.model_validate_json(f.read())
+            return State.model_validate_json(f.read())
     else:
         log(logger, MasterUninitializedLogEntry())
-        return MasterState()
+        return State()
 
 
 # FastAPI Dependencies
@@ -46,8 +46,8 @@ def check_env_vars_defined(data: object, logger: Logger) -> MasterEnvironmentSch
     return data
 
 
-def get_master_state_dependency(data: object, logger: Logger) -> MasterState:
-    if not isinstance(data, MasterState):
+def get_state_dependency(data: object, logger: Logger) -> State:
+    if not isinstance(data, State):
         raise RuntimeError("Master State Not Found")
     return data
 
diff --git a/master/placement.py b/master/placement.py
index 1d7a98fe..2eaf9ad0 100644
--- a/master/placement.py
+++ b/master/placement.py
@@ -3,7 +3,7 @@ from typing import Mapping, Sequence
 
 from shared.types.events.common import BaseEvent, EventCategory
 from shared.types.graphs.topology import Topology
-from shared.types.states.master import CachePolicy, CachePolicyType
+from shared.types.state import CachePolicy
 from shared.types.tasks.common import Task
 from shared.types.worker.instances import InstanceId, InstanceParams
 
@@ -13,7 +13,7 @@ def get_instance_placement(
     outbox: Queue[Task],
     topology: Topology,
     current_instances: Mapping[InstanceId, InstanceParams],
-    cache_policy: CachePolicy[CachePolicyType],
+    cache_policy: CachePolicy,
 ) -> Mapping[InstanceId, InstanceParams]: ...
 
 
diff --git a/shared/types/models/model.py b/shared/types/models/model.py
index faa7c3ad..c50ade27 100644
--- a/shared/types/models/model.py
+++ b/shared/types/models/model.py
@@ -8,7 +8,6 @@ from shared.types.models.sources import ModelSource
 
 
 @final
-# Concerned by the naming here; model could also be an instance of a model.
 class ModelInfo(BaseModel):
     model_id: ModelId
     model_sources: Sequence[ModelSource]
diff --git a/shared/types/state.py b/shared/types/state.py
new file mode 100644
index 00000000..59a7b1c9
--- /dev/null
+++ b/shared/types/state.py
@@ -0,0 +1,41 @@
+from collections.abc import Mapping, Sequence
+from enum import Enum
+from queue import Queue
+
+from pydantic import BaseModel, TypeAdapter
+
+from shared.types.common import NodeId
+from shared.types.graphs.topology import (
+    OrphanedPartOfTopology,
+    Topology,
+    TopologyEdge,
+    TopologyNode,
+)
+from shared.types.profiling.common import NodePerformanceProfile
+from shared.types.tasks.common import Task, TaskId, TaskSagaEntry
+from shared.types.worker.common import InstanceId, NodeStatus
+from shared.types.worker.instances import BaseInstance
+from shared.types.worker.runners import RunnerId, RunnerStatus
+
+
+class ExternalCommand(BaseModel): ...
+
+
+class CachePolicy(str, Enum):
+    KeepAll = "KeepAll"
+
+
+class State(BaseModel):
+    node_status: Mapping[NodeId, NodeStatus] = {}
+    instances: Mapping[InstanceId, BaseInstance] = {}
+    runners: Mapping[RunnerId, RunnerStatus] = {}
+    tasks: Mapping[TaskId, Task] = {}
+    task_sagas: Mapping[TaskId, Sequence[TaskSagaEntry]] = {}
+    node_profiles: Mapping[NodeId, NodePerformanceProfile] = {}
+    topology: Topology = Topology(
+        edge_base=TypeAdapter(TopologyEdge), vertex_base=TypeAdapter(TopologyNode)
+    )
+    history: Sequence[OrphanedPartOfTopology] = []
+    task_inbox: Queue[Task] = Queue()
+    task_outbox: Queue[Task] = Queue()
+    cache_policy: CachePolicy = CachePolicy.KeepAll
diff --git a/shared/types/states/master.py b/shared/types/states/master.py
deleted file mode 100644
index bb629266..00000000
--- a/shared/types/states/master.py
+++ /dev/null
@@ -1,59 +0,0 @@
-from collections.abc import Mapping, Sequence
-from enum import Enum
-from queue import Queue
-from typing import Generic, Literal, TypeVar
-
-from pydantic import BaseModel, TypeAdapter
-
-from shared.types.common import NodeId
-from shared.types.events.common import EventCategoryEnum, State
-from shared.types.graphs.topology import (
-    OrphanedPartOfTopology,
-    Topology,
-    TopologyEdge,
-    TopologyEdgeId,
-    TopologyNode,
-)
-from shared.types.profiling.common import NodePerformanceProfile
-from shared.types.states.shared import SharedState
-from shared.types.tasks.common import Task
-
-
-class ExternalCommand(BaseModel): ...
-
-
-class CachePolicyType(str, Enum):
-    KeepAll = "KeepAll"
-
-
-CachePolicyTypeT = TypeVar("CachePolicyTypeT", bound=CachePolicyType)
-
-
-class CachePolicy(BaseModel, Generic[CachePolicyTypeT]):
-    policy_type: CachePolicyTypeT
-
-
-class NodePerformanceProfileState(State[EventCategoryEnum.MutatesNodePerformanceState]):
-    node_profiles: Mapping[NodeId, NodePerformanceProfile]
-
-
-class TopologyState(State[EventCategoryEnum.MutatesTopologyState]):
-    event_category: Literal[EventCategoryEnum.MutatesTopologyState] = (
-        EventCategoryEnum.MutatesTopologyState
-    )
-    topology: Topology = Topology(
-        edge_base=TypeAdapter(TopologyEdge), vertex_base=TypeAdapter(TopologyNode)
-    )
-    history: Sequence[OrphanedPartOfTopology] = []
-
-    def delete_edge(self, edge_id: TopologyEdgeId) -> None: ...
-    def add_edge(self, edge: TopologyEdge) -> None: ...
-
-
-class MasterState(SharedState):
-    topology_state: TopologyState = TopologyState()
-    task_inbox: Queue[Task] = Queue()
-    task_outbox: Queue[Task] = Queue()
-    cache_policy: CachePolicy[CachePolicyType] = CachePolicy[CachePolicyType](
-        policy_type=CachePolicyType.KeepAll
-    )
diff --git a/shared/types/states/shared.py b/shared/types/states/shared.py
deleted file mode 100644
index ec2c06ef..00000000
--- a/shared/types/states/shared.py
+++ /dev/null
@@ -1,52 +0,0 @@
-from collections.abc import Mapping
-from typing import Literal, Sequence
-
-from pydantic import BaseModel
-
-from shared.types.common import NodeId
-from shared.types.events.common import EventCategoryEnum, State
-from shared.types.tasks.common import Task, TaskId, TaskSagaEntry
-from shared.types.worker.common import InstanceId
-from shared.types.worker.instances import BaseInstance
-from shared.types.worker.runners import RunnerId, RunnerStatus
-
-
-class Instances(State[EventCategoryEnum.MutatesInstanceState]):
-    event_category: Literal[EventCategoryEnum.MutatesInstanceState] = (
-        EventCategoryEnum.MutatesInstanceState
-    )
-    instances: Mapping[InstanceId, BaseInstance] = {}
-
-
-class Tasks(State[EventCategoryEnum.MutatesTaskState]):
-    event_category: Literal[EventCategoryEnum.MutatesTaskState] = (
-        EventCategoryEnum.MutatesTaskState
-    )
-    tasks: Mapping[TaskId, Task] = {}
-
-
-class TaskSagas(State[EventCategoryEnum.MutatesTaskSagaState]):
-    event_category: Literal[EventCategoryEnum.MutatesTaskSagaState] = (
-        EventCategoryEnum.MutatesTaskSagaState
-    )
-    task_sagas: Mapping[TaskId, Sequence[TaskSagaEntry]] = {}
-
-
-class Runners(State[EventCategoryEnum.MutatesRunnerStatus]):
-    event_category: Literal[EventCategoryEnum.MutatesRunnerStatus] = (
-        EventCategoryEnum.MutatesRunnerStatus
-    )
-    runner_statuses: Mapping[RunnerId, RunnerStatus] = {}
-
-
-class SharedState(BaseModel):
-    instances: Instances = Instances()
-    runners: Runners = Runners()
-    tasks: Tasks = Tasks()
-    task_sagas: TaskSagas = TaskSagas()
-
-    def get_node_id(self) -> NodeId: ...
-
-    def get_tasks_by_instance(
-        self, instance_id: InstanceId
-    ) -> Sequence[Task]: ...
diff --git a/shared/types/states/worker.py b/shared/types/states/worker.py
deleted file mode 100644
index 285488cf..00000000
--- a/shared/types/states/worker.py
+++ /dev/null
@@ -1,21 +0,0 @@
-from collections.abc import Mapping
-from typing import Literal
-
-from shared.types.common import NodeId
-from shared.types.events.common import (
-    EventCategoryEnum,
-    State,
-)
-from shared.types.states.shared import SharedState
-from shared.types.worker.common import NodeStatus
-
-
-class NodeStatusState(State[EventCategoryEnum.MutatesRunnerStatus]):
-    event_category: Literal[EventCategoryEnum.MutatesRunnerStatus] = (
-        EventCategoryEnum.MutatesRunnerStatus
-    )
-    node_status: Mapping[NodeId, NodeStatus]
-
-
-class WorkerState(SharedState):
-    node_status: NodeStatusState
diff --git a/worker/main.py b/worker/main.py
index 28179437..52094970 100644
--- a/worker/main.py
+++ b/worker/main.py
@@ -10,7 +10,7 @@ from pydantic import BaseModel, ConfigDict
 from shared.types.common import NodeId
 from shared.types.events.events import ChunkGenerated, InstanceId, RunnerStatusUpdated
 from shared.types.events.registry import Event
-from shared.types.states.worker import WorkerState
+from shared.types.state import State
 from shared.types.worker.common import RunnerId
 from shared.types.worker.downloads import (
     DownloadCompleted,
@@ -68,7 +68,7 @@ class Worker:
     def __init__(
         self,
         node_id: NodeId,
-        initial_state: WorkerState,
+        initial_state: State,
         logger: Logger,
     ):
         self.node_id = node_id
@@ -295,7 +295,7 @@ class Worker:
             yield event
 
     ## Planning logic
-    def plan(self, state: WorkerState) -> RunnerOp | None:
+    def plan(self, state: State) -> RunnerOp | None:
         # Compare state to worker 'mood'
         
         # First spin things down
@@ -303,7 +303,7 @@ class Worker:
         # Then spin things up
 
         # Then make sure things are downloading.
-        for instance_id, instance in state.instances.instances.items():
+        for instance_id, instance in state.instances.items():
             # We should already have asserted that this runner exists
             # If it didn't exist then we return a assign_runner op.
             for node_id, runner_id in instance.instance_params.shard_assignments.node_to_runner.items():
diff --git a/worker/test_worker_state.py b/worker/test_worker_state.py
index 5db3f9a9..99f154d7 100644
--- a/worker/test_worker_state.py
+++ b/worker/test_worker_state.py
@@ -8,7 +8,7 @@ from uuid import uuid4
 import pytest
 
 from shared.types.common import NodeId
-from shared.types.states.worker import NodeStatusState, WorkerState
+from shared.types.state import State
 from shared.types.worker.common import InstanceId, NodeStatus
 from shared.types.worker.instances import Instance
 from worker.main import Worker
@@ -31,18 +31,17 @@ async def test_worker_instance_added(worker: Worker, instance: Callable[[NodeId]
     await worker.start()
     await asyncio.sleep(0.01)
 
-    worker.state.instances.instances = {InstanceId(uuid4()): instance(worker.node_id)}
+    worker.state.instances = {InstanceId(uuid4()): instance(worker.node_id)}
     
-    print(worker.state.instances.instances)
+    print(worker.state.instances)
 
 def test_plan_noop(worker: Worker):
-    s = WorkerState(
-        node_status=NodeStatusState(
-            node_status={
+    s = State(
+        node_status={
                 NodeId(uuid4()): NodeStatus.Idle
             }
-        ),
-    )
+        )
+
     next_op = worker.plan(s)
 
     assert next_op is None
diff --git a/worker/tests/conftest.py b/worker/tests/conftest.py
index 07a67b49..afe312c0 100644
--- a/worker/tests/conftest.py
+++ b/worker/tests/conftest.py
@@ -8,7 +8,7 @@ import pytest
 
 from shared.types.common import NodeId
 from shared.types.models.common import ModelId
-from shared.types.states.worker import NodeStatusState, WorkerState
+from shared.types.state import State
 from shared.types.tasks.common import (
     ChatCompletionMessage,
     ChatCompletionTaskParams,
@@ -115,14 +115,12 @@ def chat_task(
     )
 
 @pytest.fixture
-def worker_state():
-    node_status=NodeStatusState(
-            node_status={
-                NodeId(uuid.uuid4()): NodeStatus.Idle
-            }
-        )
+def state():
+    node_status={
+        NodeId(uuid.uuid4()): NodeStatus.Idle
+    }
 
-    return WorkerState(
+    return State(
         node_status=node_status,
     )
 
@@ -157,8 +155,8 @@ def instance(pipeline_shard_meta: Callable[[int, int], PipelineShardMetadata], h
     return _instance
 
 @pytest.fixture
-def worker(worker_state: WorkerState, logger: Logger):
-    return Worker(NodeId(uuid.uuid4()), worker_state, logger)
+def worker(state: State, logger: Logger):
+    return Worker(NodeId(uuid.uuid4()), state, logger)
 
 @pytest.fixture
 async def worker_with_assigned_runner(worker: Worker, instance: Callable[[NodeId], Instance]):
diff --git a/worker/tests/test_worker_plan.py b/worker/tests/test_worker_plan.py
index 02603b85..56c0503b 100644
--- a/worker/tests/test_worker_plan.py
+++ b/worker/tests/test_worker_plan.py
@@ -9,10 +9,9 @@ import pytest
 
 from shared.types.common import NodeId
 from shared.types.models.common import ModelId
-from shared.types.states.shared import Instances
+from shared.types.state import State
 
 # WorkerState import below after RunnerCase definition to avoid forward reference issues
-from shared.types.states.worker import NodeStatusState, WorkerState
 from shared.types.worker.common import InstanceId, NodeStatus, RunnerId
 from shared.types.worker.downloads import DownloadOngoing, DownloadProgressData
 from shared.types.worker.instances import Instance, InstanceParams, TypeOfInstance
@@ -46,7 +45,7 @@ class PlanTestCase:
     expected_op_runner_idx: Optional[int] = None
     # Allow overriding the WorkerState passed to Worker.plan.  When None, a default state
     # is constructed from `runners` via helper `_build_worker_state`.
-    worker_state_override: Optional[WorkerState] = None
+    worker_state_override: Optional[State] = None
 
     def id(self) -> str:  # noqa: D401
         return self.description.replace(" ", "_")
@@ -104,9 +103,9 @@ TEST_CASES: Final[List[PlanTestCase]] = [
         ],
         expected_op_type=None,
         expected_op_runner_idx=None,
-        worker_state_override=WorkerState(
-            node_status=NodeStatusState(node_status={NodeId(): NodeStatus.Idle}),
-            instances=Instances(instances={}),
+        worker_state_override=State(
+            node_status={NodeId(): NodeStatus.Idle},
+            instances={},
         ),
     ),
 ]
@@ -130,7 +129,7 @@ def _build_worker_state(
     tmp_path: Path,
     node_id: NodeId,
     runner_cases: List[RunnerCase],
-) -> tuple[WorkerState, List[RunnerContext]]:
+) -> tuple[State, List[RunnerContext]]:
     """Construct a WorkerState plus per-runner context objects."""
 
     instances: dict[InstanceId, Instance] = {}
@@ -182,9 +181,9 @@ def _build_worker_state(
             )
         )
 
-    worker_state = WorkerState(
-        node_status=NodeStatusState(node_status={node_id: NodeStatus.Idle}),
-        instances=Instances(instances=instances),
+    worker_state = State(
+        node_status={node_id: NodeStatus.Idle},
+        instances=instances,
     )
 
     return worker_state, runner_contexts
@@ -260,4 +259,3 @@ def test_worker_plan(case: PlanTestCase, tmp_path: Path, monkeypatch: pytest.Mon
         assert op.runner_id == target_ctx.runner_id
         assert op.instance_id == target_ctx.instance_id
         assert op.shard_metadata == target_ctx.shard_metadata
-

← d19aa4f9 Simplify `Task` type + merge control & data plane types into  ·  back to Exo  ·  add forwarder supervisor 54efd01d →