[object Object]

← 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

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 →