← back to Exo
Refactor worker + master state into single state
bae58dd368486b3d829ba57f1513ddcdaf6d6ca7 · 2025-07-21 19:36:54 +0100 · Seth Howes
Files touched
M master/main.pyM master/placement.pyM shared/types/models/model.pyA shared/types/state.pyD shared/types/states/master.pyD shared/types/states/shared.pyD shared/types/states/worker.pyM worker/main.pyM worker/test_worker_state.pyM worker/tests/conftest.pyM worker/tests/test_worker_plan.py
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 →