Add Multiaddr type and refactor Hosts type for creating shard placement
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user