← back to Exo
Add Multiaddr type and refactor Hosts type for creating shard placement
e9b803604bf5421e062ea1f3d4f785c6da05aaa7 · 2025-07-28 11:39:46 +0100 · Seth Howes
Files touched
M engines/mlx/utils_mlx.pyM master/discovery_supervisor.pyM master/placement.pyM master/tests/conftest.pyM master/tests/test_placement_utils.pyM master/tests/test_topology.pyM master/utils/placement_utils.pyM shared/tests/test_state_serialization.pyM shared/topology.pyM shared/types/common.pyA shared/types/multiaddr.pyM shared/types/topology.pyM shared/types/worker/commands_runner.pyM shared/types/worker/instances.pyD shared/types/worker/mlx.pyM shared/types/worker/ops.pyM worker/main.pyM worker/runner/runner_supervisor.pyM worker/tests/conftest.pyM worker/tests/test_serdes.pyM worker/tests/test_supervisor.pyM worker/tests/test_worker_integration.py
Diff
commit e9b803604bf5421e062ea1f3d4f785c6da05aaa7
Author: Seth Howes <71157822+sethhowes@users.noreply.github.com>
Date: Mon Jul 28 11:39:46 2025 +0100
Add Multiaddr type and refactor Hosts type for creating shard placement
---
engines/mlx/utils_mlx.py | 2 +-
master/discovery_supervisor.py | 9 ++++---
master/placement.py | 11 ++++++--
master/tests/conftest.py | 13 ++++++---
master/tests/test_placement_utils.py | 37 +++++++++++++++++++++++++-
master/tests/test_topology.py | 19 ++++++++------
master/utils/placement_utils.py | 27 ++++++++++++++++++-
shared/tests/test_state_serialization.py | 5 ++--
shared/topology.py | 22 +++++++++++++++-
shared/types/common.py | 18 ++++++++++++-
shared/types/multiaddr.py | 45 ++++++++++++++++++++++++++++++++
shared/types/topology.py | 9 ++++---
shared/types/worker/commands_runner.py | 2 +-
shared/types/worker/instances.py | 2 +-
shared/types/worker/mlx.py | 17 ------------
shared/types/worker/ops.py | 2 +-
worker/main.py | 3 +--
worker/runner/runner_supervisor.py | 3 +--
worker/tests/conftest.py | 6 ++---
worker/tests/test_serdes.py | 2 +-
worker/tests/test_supervisor.py | 2 +-
worker/tests/test_worker_integration.py | 3 +--
22 files changed, 200 insertions(+), 59 deletions(-)
diff --git a/engines/mlx/utils_mlx.py b/engines/mlx/utils_mlx.py
index 781c76f9..3b7c5147 100644
--- a/engines/mlx/utils_mlx.py
+++ b/engines/mlx/utils_mlx.py
@@ -12,8 +12,8 @@ from mlx_lm.utils import load_model # type: ignore
from pydantic import RootModel
from engines.mlx.auto_parallel import auto_parallel
+from shared.types.common import Host
from shared.types.tasks import ChatCompletionTaskParams
-from shared.types.worker.mlx import Host
from shared.types.worker.shards import ShardMetadata
from worker.download.download_utils import build_model_path
from worker.runner.communication import runner_print
diff --git a/master/discovery_supervisor.py b/master/discovery_supervisor.py
index 16ed116a..440d512b 100644
--- a/master/discovery_supervisor.py
+++ b/master/discovery_supervisor.py
@@ -6,6 +6,7 @@ from exo_pyo3_bindings import ConnectionUpdate, DiscoveryService, Keypair
from shared.db import AsyncSQLiteEventStorage
from shared.types.common import NodeId
from shared.types.events import TopologyEdgeCreated, TopologyEdgeDeleted
+from shared.types.multiaddr import Multiaddr
from shared.types.topology import Connection
@@ -44,8 +45,8 @@ class DiscoverySupervisor:
async def _connected_callback(self, e: ConnectionUpdate) -> None:
local_node_id = self.node_id
send_back_node_id = NodeId(e.peer_id.to_base58())
- local_multiaddr = e.local_addr.to_string()
- send_back_multiaddr = e.send_back_addr.to_string()
+ local_multiaddr = Multiaddr(address=str(e.local_addr))
+ send_back_multiaddr = Multiaddr(address=str(e.send_back_addr))
connection_profile = None
topology_edge_created = TopologyEdgeCreated(edge=Connection(
@@ -65,8 +66,8 @@ class DiscoverySupervisor:
async def _disconnected_callback(self, e: ConnectionUpdate) -> None:
local_node_id = self.node_id
send_back_node_id = NodeId(e.peer_id.to_base58())
- local_multiaddr = e.local_addr.to_string()
- send_back_multiaddr = e.send_back_addr.to_string()
+ local_multiaddr = Multiaddr(address=str(e.local_addr))
+ send_back_multiaddr = Multiaddr(address=str(e.send_back_addr))
connection_profile = None
topology_edge_created = TopologyEdgeDeleted(edge=Connection(
diff --git a/master/placement.py b/master/placement.py
index cd3320cc..e502f5d3 100644
--- a/master/placement.py
+++ b/master/placement.py
@@ -1,4 +1,3 @@
-
from collections.abc import Mapping
from copy import deepcopy
from functools import singledispatch
@@ -6,10 +5,12 @@ from typing import Sequence
from master.utils.placement_utils import (
filter_cycles_by_memory,
+ get_hosts_from_subgraph,
get_shard_assignments,
get_smallest_cycles,
)
from shared.topology import Topology
+from shared.types.common import Host
from shared.types.events import Event, InstanceCreated, InstanceDeleted
from shared.types.events.commands import CreateInstanceCommand, DeleteInstanceCommand
from shared.types.worker.common import InstanceId
@@ -40,13 +41,19 @@ def get_instance_placements(
shard_assignments = get_shard_assignments(command.model_meta, selected_cycle)
+ cycle_digraph: Topology = topology.get_subgraph_from_nodes(selected_cycle)
+ hosts: list[Host] = get_hosts_from_subgraph(cycle_digraph)
+
instance_id = command.instance_id
target_instances = deepcopy(current_instances)
target_instances[instance_id] = Instance(
instance_id=instance_id,
instance_type=InstanceStatus.ACTIVE,
shard_assignments=shard_assignments,
- hosts=[]
+ hosts=[Host(
+ ip=host.ip,
+ port=host.port,
+ ) for host in hosts]
)
return target_instances
diff --git a/master/tests/conftest.py b/master/tests/conftest.py
index 6aee767a..1fbabfc8 100644
--- a/master/tests/conftest.py
+++ b/master/tests/conftest.py
@@ -1,6 +1,7 @@
import pytest
from shared.types.common import NodeId
+from shared.types.multiaddr import Multiaddr
from shared.types.profiling import (
MemoryPerformanceProfile,
NodePerformanceProfile,
@@ -33,14 +34,20 @@ def create_node():
return _create_node
+# TODO: this is a hack to get the port for the send_back_multiaddr
@pytest.fixture
def create_connection():
- def _create_connection(source_node_id: NodeId, sink_node_id: NodeId) -> Connection:
+ port_counter = 1235
+ def _create_connection(source_node_id: NodeId, sink_node_id: NodeId, send_back_port: int | None = None) -> Connection:
+ nonlocal port_counter
+ if send_back_port is None:
+ send_back_port = port_counter
+ port_counter += 1
return Connection(
local_node_id=source_node_id,
send_back_node_id=sink_node_id,
- local_multiaddr="/ip4/127.0.0.1/tcp/1234",
- send_back_multiaddr="/ip4/127.0.0.1/tcp/1235",
+ local_multiaddr=Multiaddr(address="/ip4/127.0.0.1/tcp/1234"),
+ send_back_multiaddr=Multiaddr(address=f"/ip4/127.0.0.1/tcp/{send_back_port}"),
connection_profile=ConnectionProfile(throughput=1000, latency=1000, jitter=1000)
)
diff --git a/master/tests/test_placement_utils.py b/master/tests/test_placement_utils.py
index d898f89a..646aa994 100644
--- a/master/tests/test_placement_utils.py
+++ b/master/tests/test_placement_utils.py
@@ -1,14 +1,16 @@
+from ipaddress import IPv4Address
from typing import Callable
import pytest
from master.utils.placement_utils import (
filter_cycles_by_memory,
+ get_hosts_from_subgraph,
get_shard_assignments,
get_smallest_cycles,
)
from shared.topology import Topology
-from shared.types.common import NodeId
+from shared.types.common import Host, NodeId
from shared.types.models import ModelMetadata
from shared.types.topology import Connection, Node
@@ -173,3 +175,36 @@ def test_get_shard_assignments(topology: Topology, create_node: Callable[[int, N
assert shard_assignments.runner_to_shard[runner_id_c].end_layer - shard_assignments.runner_to_shard[runner_id_c].start_layer == expected_layers[2]
assert shard_assignments.runner_to_shard[runner_id_a].end_layer - shard_assignments.runner_to_shard[runner_id_a].start_layer == expected_layers[0]
assert shard_assignments.runner_to_shard[runner_id_b].end_layer - shard_assignments.runner_to_shard[runner_id_b].start_layer == expected_layers[1]
+
+
+def test_get_hosts_from_subgraph(topology: Topology, create_node: Callable[[int, NodeId | None], Node], create_connection: Callable[[NodeId, NodeId, int | None], Connection]):
+ # arrange
+ node_a_id = NodeId()
+ node_b_id = NodeId()
+ node_c_id = NodeId()
+
+ node_a = create_node(500, node_a_id)
+ node_b = create_node(500, node_b_id)
+ node_c = create_node(1000, node_c_id)
+
+ topology.add_node(node_a)
+ topology.add_node(node_b)
+ topology.add_node(node_c)
+
+ topology.add_connection(create_connection(node_a_id, node_b_id, 5001))
+ topology.add_connection(create_connection(node_b_id, node_c_id, 5002))
+ topology.add_connection(create_connection(node_c_id, node_a_id, 5003))
+ topology.add_connection(create_connection(node_b_id, node_a_id, 5004))
+
+ # act
+ hosts = get_hosts_from_subgraph(topology)
+
+ # assert
+ assert len(hosts) == 3
+ expected_hosts = [
+ Host(ip=IPv4Address("127.0.0.1"), port=5001),
+ Host(ip=IPv4Address("127.0.0.1"), port=5002),
+ Host(ip=IPv4Address("127.0.0.1"), port=5003),
+ ]
+ for expected_host in expected_hosts:
+ assert expected_host in hosts
diff --git a/master/tests/test_topology.py b/master/tests/test_topology.py
index 5264c7b6..9765c20d 100644
--- a/master/tests/test_topology.py
+++ b/master/tests/test_topology.py
@@ -1,6 +1,7 @@
import pytest
from shared.topology import Topology
+from shared.types.multiaddr import Multiaddr
from shared.types.profiling import (
MemoryPerformanceProfile,
NodePerformanceProfile,
@@ -16,10 +17,12 @@ def topology() -> Topology:
@pytest.fixture
def connection() -> Connection:
- return Connection(local_node_id=NodeId(), send_back_node_id=NodeId(), local_multiaddr="/ip4/127.0.0.1/tcp/1234",
- send_back_multiaddr="/ip4/127.0.0.1/tcp/1235",
- connection_profile=ConnectionProfile(throughput=1000, latency=1000, jitter=1000))
-
+ return Connection(
+ local_node_id=NodeId(),
+ send_back_node_id=NodeId(),
+ local_multiaddr=Multiaddr(address="/ip4/127.0.0.1/tcp/1234"),
+ send_back_multiaddr=Multiaddr(address="/ip4/127.0.0.1/tcp/1235"),
+ connection_profile=ConnectionProfile(throughput=1000, latency=1000, jitter=1000))
@pytest.fixture
def node_profile() -> NodePerformanceProfile:
@@ -128,16 +131,16 @@ def test_remove_connection_bridge(topology: Topology, node_profile: NodePerforma
connection_master_to_a = Connection(
local_node_id=master_id,
send_back_node_id=node_a_id,
- local_multiaddr="/ip4/127.0.0.1/tcp/1234",
- send_back_multiaddr="/ip4/127.0.0.1/tcp/1235",
+ local_multiaddr=Multiaddr(address="/ip4/127.0.0.1/tcp/1234"),
+ send_back_multiaddr=Multiaddr(address="/ip4/127.0.0.1/tcp/1235"),
connection_profile=ConnectionProfile(throughput=1000, latency=1000, jitter=1000)
)
connection_a_to_b = Connection(
local_node_id=node_a_id,
send_back_node_id=node_b_id,
- local_multiaddr="/ip4/127.0.0.1/tcp/1236",
- send_back_multiaddr="/ip4/127.0.0.1/tcp/1237",
+ local_multiaddr=Multiaddr(address="/ip4/127.0.0.1/tcp/1236"),
+ send_back_multiaddr=Multiaddr(address="/ip4/127.0.0.1/tcp/1237"),
connection_profile=ConnectionProfile(throughput=1000, latency=1000, jitter=1000)
)
diff --git a/master/utils/placement_utils.py b/master/utils/placement_utils.py
index 30d96725..157f2182 100644
--- a/master/utils/placement_utils.py
+++ b/master/utils/placement_utils.py
@@ -2,7 +2,8 @@ from typing import TypeGuard, cast
from pydantic import BaseModel
-from shared.types.common import NodeId
+from shared.topology import Topology
+from shared.types.common import Host, NodeId
from shared.types.models import ModelMetadata
from shared.types.profiling import NodePerformanceProfile
from shared.types.topology import Node
@@ -75,3 +76,27 @@ def get_shard_assignments(
)
return shard_assignments
+
+
+def get_hosts_from_subgraph(cycle_digraph: Topology) -> list[Host]:
+ cycles = cycle_digraph.get_cycles()
+ if not cycles:
+ return []
+
+ cycle = cycles[0]
+ hosts: list[Host] = []
+ for i in range(len(cycle)):
+ current_node = cycle[i]
+ next_node = cycle[(i + 1) % len(cycle)]
+
+ for connection in cycle_digraph.list_connections():
+ if (connection.local_node_id == current_node.node_id and
+ connection.send_back_node_id == next_node.node_id):
+ host = Host(
+ ip=connection.send_back_multiaddr.ipv4_address,
+ port=connection.send_back_multiaddr.port
+ )
+ hosts.append(host)
+ break
+
+ return hosts
\ No newline at end of file
diff --git a/shared/tests/test_state_serialization.py b/shared/tests/test_state_serialization.py
index c41e0cc3..35d42c1e 100644
--- a/shared/tests/test_state_serialization.py
+++ b/shared/tests/test_state_serialization.py
@@ -1,6 +1,7 @@
from __future__ import annotations
from shared.types.common import NodeId
+from shared.types.multiaddr import Multiaddr
from shared.types.state import State
from shared.types.topology import Connection
@@ -15,8 +16,8 @@ def test_state_serialization_roundtrip() -> None:
connection = Connection(
local_node_id=node_a,
send_back_node_id=node_b,
- local_multiaddr="/ip4/127.0.0.1/tcp/10000",
- send_back_multiaddr="/ip4/127.0.0.1/tcp/10001",
+ local_multiaddr=Multiaddr(address="/ip4/127.0.0.1/tcp/10000"),
+ send_back_multiaddr=Multiaddr(address="/ip4/127.0.0.1/tcp/10001"),
)
state = State()
diff --git a/shared/topology.py b/shared/topology.py
index 2263c447..cdbc6622 100644
--- a/shared/topology.py
+++ b/shared/topology.py
@@ -5,6 +5,7 @@ import rustworkx as rx
from pydantic import BaseModel, ConfigDict
from shared.types.common import NodeId
+from shared.types.multiaddr import Multiaddr
from shared.types.profiling import ConnectionProfile, NodePerformanceProfile
from shared.types.topology import Connection, Node, TopologyProto
@@ -94,7 +95,15 @@ class Topology(TopologyProto):
def get_node_profile(self, node_id: NodeId) -> NodePerformanceProfile | None:
rx_idx = self._node_id_to_rx_id_map[node_id]
return self._graph.get_node_data(rx_idx).node_profile
-
+
+ def get_node_multiaddr(self, node_id: NodeId) -> Multiaddr:
+ for connection in self.list_connections():
+ if connection.local_node_id == node_id:
+ return connection.local_multiaddr
+ if connection.send_back_node_id == node_id:
+ return connection.send_back_multiaddr
+ raise ValueError(f"Node {node_id} is not connected to any other nodes")
+
def update_node_profile(self, node_id: NodeId, node_profile: NodePerformanceProfile) -> None:
rx_idx = self._node_id_to_rx_id_map[node_id]
self._graph[rx_idx].node_profile = node_profile
@@ -137,6 +146,17 @@ class Topology(TopologyProto):
cycles.append(cycle)
return cycles
+
+ def get_subgraph_from_nodes(self, nodes: list[Node]) -> "Topology":
+ node_idxs = [node.node_id for node in nodes]
+ rx_idxs = [self._node_id_to_rx_id_map[idx] for idx in node_idxs]
+ topology = Topology()
+ for rx_idx in rx_idxs:
+ topology.add_node(self._graph[rx_idx])
+ for connection in self.list_connections():
+ if connection.local_node_id in node_idxs and connection.send_back_node_id in node_idxs:
+ topology.add_connection(connection)
+ return topology
def _is_bridge(self, connection: Connection) -> bool:
edge_idx = self._edge_id_to_rx_id_map[connection]
diff --git a/shared/types/common.py b/shared/types/common.py
index 0cd167ab..a5e441a3 100644
--- a/shared/types/common.py
+++ b/shared/types/common.py
@@ -1,7 +1,8 @@
+from ipaddress import IPv4Address
from typing import Any, Self
from uuid import uuid4
-from pydantic import GetCoreSchemaHandler
+from pydantic import BaseModel, GetCoreSchemaHandler, field_validator
from pydantic_core import core_schema
@@ -25,3 +26,18 @@ class NodeId(ID):
class CommandId(ID):
pass
+
+
+class Host(BaseModel):
+ ip: IPv4Address
+ port: int
+
+ def __str__(self) -> str:
+ return f"{self.ip}:{self.port}"
+
+ @field_validator("port")
+ @classmethod
+ def check_port(cls, v: int) -> int:
+ if not (0 <= v <= 65535):
+ raise ValueError("Port must be between 0 and 65535")
+ return v
diff --git a/shared/types/multiaddr.py b/shared/types/multiaddr.py
new file mode 100644
index 00000000..53c0a22f
--- /dev/null
+++ b/shared/types/multiaddr.py
@@ -0,0 +1,45 @@
+import re
+from ipaddress import IPv4Address
+from typing import ClassVar
+
+from pydantic import BaseModel, computed_field, field_validator
+
+
+class Multiaddr(BaseModel):
+ address: str
+
+ PATTERNS: ClassVar[list[str]] = [
+ r'^/ip4/(\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3})(/tcp/(\d{1,5}))?(/p2p/[A-Za-z0-9]+)?$',
+ r'^/ip6/([0-9a-fA-F:]+)(/tcp/(\d{1,5}))?(/p2p/[A-Za-z0-9]+)?$',
+ r'^/dns[46]?/([a-zA-Z0-9.-]+)(/tcp/(\d{1,5}))?(/p2p/[A-Za-z0-9]+)?$',
+ ]
+
+ @field_validator("address")
+ @classmethod
+ def validate_format(cls, v: str) -> str:
+ if not any(re.match(pattern, v) for pattern in cls.PATTERNS):
+ raise ValueError(
+ f"Invalid multiaddr format: {v}. "
+ "Expected format like /ip4/127.0.0.1/tcp/4001 or /dns/example.com/tcp/443"
+ )
+ return v
+
+ @computed_field
+ @property
+ def ipv4_address(self) -> IPv4Address:
+ match = re.match(r'^/ip4/(\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3})', self.address)
+ if not match:
+ raise ValueError(f"Invalid multiaddr format: {self.address}. Expected format like /ip4/127.0.0.1/tcp/4001")
+ return IPv4Address(match.group(1))
+
+ @computed_field
+ @property
+ def port(self) -> int:
+ match = re.search(r'/tcp/(\d{1,5})', self.address)
+ if not match:
+ raise ValueError(f"Invalid multiaddr format: {self.address}. Expected format like /ip4/127.0.0.1/tcp/4001")
+ return int(match.group(1))
+
+
+ def __str__(self) -> str:
+ return self.address
diff --git a/shared/types/topology.py b/shared/types/topology.py
index de32abd1..2a5609fd 100644
--- a/shared/types/topology.py
+++ b/shared/types/topology.py
@@ -3,14 +3,15 @@ from typing import Iterable, Protocol
from pydantic import BaseModel, ConfigDict
from shared.types.common import NodeId
+from shared.types.multiaddr import Multiaddr
from shared.types.profiling import ConnectionProfile, NodePerformanceProfile
class Connection(BaseModel):
local_node_id: NodeId
send_back_node_id: NodeId
- local_multiaddr: str
- send_back_multiaddr: str
+ local_multiaddr: Multiaddr
+ send_back_multiaddr: Multiaddr
connection_profile: ConnectionProfile | None = None
# required for Connection to be used as a key
@@ -21,8 +22,8 @@ class Connection(BaseModel):
(
self.local_node_id,
self.send_back_node_id,
- self.local_multiaddr,
- self.send_back_multiaddr,
+ self.local_multiaddr.address,
+ self.send_back_multiaddr.address,
)
)
diff --git a/shared/types/worker/commands_runner.py b/shared/types/worker/commands_runner.py
index 4432b6d7..4a05b09b 100644
--- a/shared/types/worker/commands_runner.py
+++ b/shared/types/worker/commands_runner.py
@@ -4,8 +4,8 @@ from typing import Annotated, Generic, Literal, TypeVar
from pydantic import BaseModel, Field, TypeAdapter
from shared.openai_compat import FinishReason
+from shared.types.common import Host
from shared.types.tasks import ChatCompletionTaskParams
-from shared.types.worker.mlx import Host
from shared.types.worker.shards import ShardMetadata
## Messages passed TO the runner
diff --git a/shared/types/worker/instances.py b/shared/types/worker/instances.py
index 4bfa92af..61961afc 100644
--- a/shared/types/worker/instances.py
+++ b/shared/types/worker/instances.py
@@ -2,8 +2,8 @@ from enum import Enum
from pydantic import BaseModel
+from shared.types.common import Host
from shared.types.worker.common import InstanceId
-from shared.types.worker.mlx import Host
from shared.types.worker.runners import (
ShardAssignments,
)
diff --git a/shared/types/worker/mlx.py b/shared/types/worker/mlx.py
deleted file mode 100644
index 9e8267bc..00000000
--- a/shared/types/worker/mlx.py
+++ /dev/null
@@ -1,17 +0,0 @@
-from pydantic import BaseModel, field_validator
-
-
-# TODO: Is this the right place for this? Host is consumed by worker, but typically stored in the master
-class Host(BaseModel):
- host: str
- port: int
-
- def __str__(self) -> str:
- return f"{self.host}:{self.port}"
-
- @field_validator("port")
- @classmethod
- def check_port(cls, v: int) -> int:
- if not (0 <= v <= 65535):
- raise ValueError("Port must be between 0 and 65535")
- return v
diff --git a/shared/types/worker/ops.py b/shared/types/worker/ops.py
index 97787fba..82db7c77 100644
--- a/shared/types/worker/ops.py
+++ b/shared/types/worker/ops.py
@@ -3,10 +3,10 @@ from typing import Annotated, Generic, Literal, TypeVar, Union
from pydantic import BaseModel, Field
+from shared.types.common import Host
from shared.types.events import InstanceId
from shared.types.tasks import Task
from shared.types.worker.common import RunnerId
-from shared.types.worker.mlx import Host
from shared.types.worker.shards import ShardMetadata
diff --git a/worker/main.py b/worker/main.py
index 1275a3e6..4c40d826 100644
--- a/worker/main.py
+++ b/worker/main.py
@@ -11,7 +11,7 @@ from pydantic import BaseModel, ConfigDict
from shared.apply import apply
from shared.db.sqlite import AsyncSQLiteEventStorage
from shared.db.sqlite.event_log_manager import EventLogConfig, EventLogManager
-from shared.types.common import NodeId
+from shared.types.common import Host, NodeId
from shared.types.events import (
ChunkGenerated,
Event,
@@ -32,7 +32,6 @@ from shared.types.worker.downloads import (
DownloadProgressData,
)
from shared.types.worker.instances import InstanceStatus
-from shared.types.worker.mlx import Host
from shared.types.worker.ops import (
AssignRunnerOp,
DownloadOp,
diff --git a/worker/runner/runner_supervisor.py b/worker/runner/runner_supervisor.py
index 43b515dc..8d813697 100644
--- a/worker/runner/runner_supervisor.py
+++ b/worker/runner/runner_supervisor.py
@@ -5,7 +5,7 @@ from collections.abc import AsyncGenerator
from types import CoroutineType
from typing import Any, Callable
-from shared.types.common import CommandId
+from shared.types.common import CommandId, Host
from shared.types.events.chunks import GenerationChunk, TokenChunk
from shared.types.tasks import ChatCompletionTaskParams, Task
from shared.types.worker.commands_runner import (
@@ -18,7 +18,6 @@ from shared.types.worker.commands_runner import (
RunnerResponse,
SetupMessage,
)
-from shared.types.worker.mlx import Host
from shared.types.worker.shards import ShardMetadata
from worker.runner.communication import (
supervisor_read_response,
diff --git a/worker/tests/conftest.py b/worker/tests/conftest.py
index 9ef65c3d..1808323b 100644
--- a/worker/tests/conftest.py
+++ b/worker/tests/conftest.py
@@ -1,4 +1,5 @@
import asyncio
+from ipaddress import IPv4Address
from logging import Logger, getLogger
from pathlib import Path
from typing import Awaitable, Callable
@@ -9,7 +10,7 @@ from shared.db.sqlite.connector import AsyncSQLiteEventStorage
from shared.db.sqlite.event_log_manager import EventLogConfig, EventLogManager
from shared.models.model_meta import get_model_meta
from shared.types.api import ChatCompletionMessage, ChatCompletionTaskParams
-from shared.types.common import CommandId, NodeId
+from shared.types.common import CommandId, Host, NodeId
from shared.types.models import ModelId, ModelMetadata
from shared.types.state import State
from shared.types.tasks import (
@@ -20,7 +21,6 @@ from shared.types.tasks import (
)
from shared.types.worker.common import InstanceId, NodeStatus
from shared.types.worker.instances import Instance, InstanceStatus
-from shared.types.worker.mlx import Host
from shared.types.worker.ops import (
AssignRunnerOp,
RunnerUpOp,
@@ -36,7 +36,7 @@ def hosts():
def _hosts(count: int, offset: int = 0) -> list[Host]:
return [
Host(
- host="127.0.0.1",
+ ip=IPv4Address("127.0.0.1"),
port=5000 + offset + i,
)
for i in range(count)
diff --git a/worker/tests/test_serdes.py b/worker/tests/test_serdes.py
index 42af427e..37fe515a 100644
--- a/worker/tests/test_serdes.py
+++ b/worker/tests/test_serdes.py
@@ -3,6 +3,7 @@ from typing import Callable, TypeVar
from pydantic import BaseModel, TypeAdapter
+from shared.types.common import Host
from shared.types.tasks import Task
from shared.types.worker.commands_runner import (
ChatTaskMessage,
@@ -10,7 +11,6 @@ from shared.types.worker.commands_runner import (
SetupMessage,
)
from shared.types.worker.common import InstanceId
-from shared.types.worker.mlx import Host
from shared.types.worker.shards import PipelineShardMetadata
T = TypeVar("T", bound=BaseModel)
diff --git a/worker/tests/test_supervisor.py b/worker/tests/test_supervisor.py
index 5a77eccd..77cebdf1 100644
--- a/worker/tests/test_supervisor.py
+++ b/worker/tests/test_supervisor.py
@@ -5,6 +5,7 @@ from typing import Callable
import pytest
from shared.openai_compat import FinishReason
+from shared.types.common import Host
from shared.types.events.chunks import TokenChunk
from shared.types.tasks import (
ChatCompletionTaskParams,
@@ -12,7 +13,6 @@ from shared.types.tasks import (
TaskType,
)
from shared.types.worker.common import InstanceId
-from shared.types.worker.mlx import Host
from shared.types.worker.shards import PipelineShardMetadata
from worker.runner.runner_supervisor import RunnerSupervisor
diff --git a/worker/tests/test_worker_integration.py b/worker/tests/test_worker_integration.py
index acd28735..cbd6a681 100644
--- a/worker/tests/test_worker_integration.py
+++ b/worker/tests/test_worker_integration.py
@@ -7,7 +7,7 @@ import pytest
# TaskStateUpdated and ChunkGenerated are used in test_worker_integration_utils.py
from shared.db.sqlite.connector import AsyncSQLiteEventStorage
from shared.db.sqlite.event_log_manager import EventLogConfig, EventLogManager
-from shared.types.common import NodeId
+from shared.types.common import Host, NodeId
from shared.types.events import (
InstanceCreated,
InstanceDeleted,
@@ -24,7 +24,6 @@ from shared.types.worker.instances import (
InstanceStatus,
ShardAssignments,
)
-from shared.types.worker.mlx import Host
from shared.types.worker.runners import (
FailedRunnerStatus,
LoadedRunnerStatus,
← b285a9f0 fix placement tests
·
back to Exo
·
Fix download tests 36a5d75e →