Add Multiaddr type and refactor Hosts type for creating shard placement

This commit is contained in:
Seth Howes
2025-07-28 11:39:46 +01:00
committed by GitHub
parent b285a9f0b7
commit e9b803604b
22 changed files with 200 additions and 59 deletions
+1 -1
View File
@@ -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
+5 -4
View File
@@ -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(
+9 -2
View File
@@ -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
+10 -3
View File
@@ -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)
)
+36 -1
View File
@@ -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
+11 -8
View File
@@ -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)
)
+26 -1
View File
@@ -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
+3 -2
View File
@@ -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()
+21 -1
View File
@@ -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]
+17 -1
View File
@@ -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
+45
View File
@@ -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
+5 -4
View File
@@ -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,
)
)
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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,
)
-17
View File
@@ -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
+1 -1
View File
@@ -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
+1 -2
View File
@@ -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,
+1 -2
View File
@@ -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,
+3 -3
View File
@@ -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)
+1 -1
View File
@@ -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)
+1 -1
View File
@@ -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
+1 -2
View File
@@ -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,