← back to Exo
fix: master tests
35ab6b376e51d746ad9c61c9ed190aa01cac2023 · 2025-10-07 15:36:05 +0100 · Alex Cheema
Co-authored-by: Evan <evanev7@gmail.com>
Files touched
M src/exo/master/api.pyM src/exo/master/tests/api_utils_test.pyM src/exo/master/tests/conftest.pyM src/exo/master/tests/test_master.pyM src/exo/master/tests/test_placement.pyM src/exo/master/tests/test_placement_utils.pyM src/exo/utils/channels.py
Diff
commit 35ab6b376e51d746ad9c61c9ed190aa01cac2023
Author: Alex Cheema <41707476+AlexCheema@users.noreply.github.com>
Date: Tue Oct 7 15:36:05 2025 +0100
fix: master tests
Co-authored-by: Evan <evanev7@gmail.com>
---
src/exo/master/api.py | 3 +-
src/exo/master/tests/api_utils_test.py | 5 +-
src/exo/master/tests/conftest.py | 2 +-
src/exo/master/tests/test_master.py | 210 +++++++++++++--------------
src/exo/master/tests/test_placement.py | 8 +-
src/exo/master/tests/test_placement_utils.py | 6 +-
src/exo/utils/channels.py | 2 +-
7 files changed, 118 insertions(+), 118 deletions(-)
diff --git a/src/exo/master/api.py b/src/exo/master/api.py
index 83ef17a5..d10f7dd6 100644
--- a/src/exo/master/api.py
+++ b/src/exo/master/api.py
@@ -260,7 +260,8 @@ class API:
)
await self._send(command)
return StreamingResponse(
- self._generate_chat_stream(command.command_id), media_type="text/event-stream"
+ self._generate_chat_stream(command.command_id),
+ media_type="text/event-stream",
)
def _calculate_total_available_memory(self) -> int:
diff --git a/src/exo/master/tests/api_utils_test.py b/src/exo/master/tests/api_utils_test.py
index 5682f0e5..3ed52c7a 100644
--- a/src/exo/master/tests/api_utils_test.py
+++ b/src/exo/master/tests/api_utils_test.py
@@ -19,7 +19,7 @@ from openai.types.chat import (
)
from openai.types.chat.chat_completion_chunk import ChatCompletionChunk, Choice
-from exo.master.main import async_main as master_main
+from exo.main import main
_P = ParamSpec("_P")
_R = TypeVar("_R")
@@ -34,7 +34,8 @@ def with_master_main(
@pytest.mark.asyncio
@functools.wraps(func)
async def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _R:
- master_task = asyncio.create_task(master_main())
+ loop = asyncio.get_running_loop()
+ master_task = loop.run_in_executor(None, main)
try:
return await func(*args, **kwargs)
finally:
diff --git a/src/exo/master/tests/conftest.py b/src/exo/master/tests/conftest.py
index a22333b9..39aa2b31 100644
--- a/src/exo/master/tests/conftest.py
+++ b/src/exo/master/tests/conftest.py
@@ -53,7 +53,7 @@ def create_connection() -> Callable[[NodeId, NodeId, int | None], Connection]:
local_node_id=source_node_id,
send_back_node_id=sink_node_id,
send_back_multiaddr=Multiaddr(
- address=f"/ip4/127.0.0.1/tcp/{send_back_port}"
+ address=f"/ip4/169.254.0.1/tcp/{send_back_port}"
),
connection_profile=ConnectionProfile(
throughput=1000, latency=1000, jitter=1000
diff --git a/src/exo/master/tests/test_master.py b/src/exo/master/tests/test_master.py
index b93f2bb7..bfa3f564 100644
--- a/src/exo/master/tests/test_master.py
+++ b/src/exo/master/tests/test_master.py
@@ -1,162 +1,160 @@
import asyncio
-import tempfile
-from pathlib import Path
from typing import List, Sequence
import pytest
from exo.master.main import Master
-from exo.shared.db.config import EventLogConfig
-from exo.shared.db.connector import AsyncSQLiteEventStorage
-from exo.shared.db.event_log_manager import EventLogManager
-from exo.shared.keypair import Keypair
+from exo.routing.router import get_node_id_keypair
from exo.shared.types.api import ChatCompletionMessage, ChatCompletionTaskParams
from exo.shared.types.commands import (
ChatCompletion,
- Command,
CommandId,
CreateInstance,
+ ForwarderCommand,
+ TaggedCommand,
)
from exo.shared.types.common import NodeId
from exo.shared.types.events import (
+ ForwarderEvent,
IndexedEvent,
InstanceCreated,
NodePerformanceMeasured,
+ TaggedEvent,
TaskCreated,
- TopologyNodeCreated,
)
-from exo.shared.types.models import ModelMetadata
+from exo.shared.types.memory import Memory
+from exo.shared.types.models import ModelId, ModelMetadata
from exo.shared.types.profiling import (
MemoryPerformanceProfile,
NodePerformanceProfile,
SystemPerformanceProfile,
)
from exo.shared.types.tasks import ChatCompletionTask, TaskStatus, TaskType
-from exo.shared.types.worker.instances import (
- Instance,
- InstanceStatus,
- ShardAssignments,
-)
+from exo.shared.types.worker.instances import Instance, InstanceStatus, ShardAssignments
from exo.shared.types.worker.shards import PartitionStrategy, PipelineShardMetadata
-
-
-def _create_forwarder_dummy_binary() -> Path:
- path = Path(tempfile.mkstemp()[1]) / "forwarder.bin"
- if not path.exists():
- path.parent.mkdir(parents=True, exist_ok=True)
- path.write_bytes(b"#!/bin/sh\necho dummy forwarder && sleep 1000000\n")
- path.chmod(0o755)
- return path
+from exo.utils.channels import channel
@pytest.mark.asyncio
async def test_master():
- event_log_manager = EventLogManager(EventLogConfig())
- await event_log_manager.initialize()
- global_events: AsyncSQLiteEventStorage = event_log_manager.global_events
- await global_events.delete_all_events()
+ keypair = get_node_id_keypair()
+ node_id = NodeId(keypair.to_peer_id().to_base58())
+
+ ge_sender, global_event_receiver = channel[ForwarderEvent]()
+ command_sender, co_receiver = channel[ForwarderCommand]()
+ local_event_sender, le_receiver = channel[ForwarderEvent]()
+
+ all_events: List[IndexedEvent] = []
async def _get_events() -> Sequence[IndexedEvent]:
- orig_events = await global_events.get_events_since(0)
- override_idx_in_log = 1
- events: List[IndexedEvent] = []
+ orig_events = global_event_receiver.collect()
for e in orig_events:
- events.append(
+ all_events.append(
IndexedEvent(
- event=e.event,
- idx=override_idx_in_log, # origin=e.origin,
+ event=e.tagged_event.c,
+ idx=len(all_events), # origin=e.origin,
)
)
- override_idx_in_log += 1
- return events
-
- command_buffer: List[Command] = []
+ return all_events
- forwarder_binary_path = _create_forwarder_dummy_binary()
-
- node_id_keypair = Keypair.generate_ed25519()
- node_id = NodeId(node_id_keypair.to_peer_id().to_base58())
master = Master(
- node_id_keypair,
node_id,
- command_buffer=command_buffer,
- global_events=global_events,
- forwarder_binary_path=forwarder_binary_path,
- worker_events=global_events,
+ global_event_sender=ge_sender,
+ local_event_receiver=le_receiver,
+ command_receiver=co_receiver,
+ tb_only=False,
)
asyncio.create_task(master.run())
- # wait for initial topology event
- while len(list(master.state.topology.list_nodes())) == 0:
- print("waiting")
- await asyncio.sleep(0.001)
+
+ sender_node_id = NodeId(f"{keypair.to_peer_id().to_base58()}_sender")
# inject a NodePerformanceProfile event
- await event_log_manager.global_events.append_events(
- [
- NodePerformanceMeasured(
- node_id=node_id,
- node_profile=NodePerformanceProfile(
- model_id="maccy",
- chip_id="arm",
- friendly_name="test",
- memory=MemoryPerformanceProfile(
- ram_total=678948 * 1024,
- ram_available=678948 * 1024,
- swap_total=0,
- swap_available=0,
+ await local_event_sender.send(
+ ForwarderEvent(
+ origin_idx=0,
+ origin=sender_node_id,
+ tagged_event=TaggedEvent.from_(
+ NodePerformanceMeasured(
+ node_id=node_id,
+ node_profile=NodePerformanceProfile(
+ model_id="maccy",
+ chip_id="arm",
+ friendly_name="test",
+ memory=MemoryPerformanceProfile(
+ ram_total=Memory.from_bytes(678948 * 1024),
+ ram_available=Memory.from_bytes(678948 * 1024),
+ swap_total=Memory.from_bytes(0),
+ swap_available=Memory.from_bytes(0),
+ ),
+ network_interfaces=[],
+ system=SystemPerformanceProfile(flops_fp16=0),
),
- network_interfaces=[],
- system=SystemPerformanceProfile(flops_fp16=0),
- ),
- )
- ],
- origin=node_id,
+ )
+ ),
+ )
)
+
+ # wait for initial topology event
+ while len(list(master.state.topology.list_nodes())) == 0:
+ await asyncio.sleep(0.001)
while len(master.state.node_profiles) == 0:
await asyncio.sleep(0.001)
- command_buffer.append(
- CreateInstance(
- command_id=CommandId(),
- model_meta=ModelMetadata(
- model_id="llama-3.2-1b",
- pretty_name="Llama 3.2 1B",
- n_layers=16,
- storage_size_kilobytes=678948,
+ await command_sender.send(
+ ForwarderCommand(
+ origin=node_id,
+ tagged_command=TaggedCommand.from_(
+ CreateInstance(
+ command_id=CommandId(),
+ model_meta=ModelMetadata(
+ model_id=ModelId("llama-3.2-1b"),
+ pretty_name="Llama 3.2 1B",
+ n_layers=16,
+ storage_size=Memory.from_bytes(678948),
+ ),
+ )
),
)
)
while len(master.state.instances.keys()) == 0:
await asyncio.sleep(0.001)
- command_buffer.append(
- ChatCompletion(
- command_id=CommandId(),
- request_params=ChatCompletionTaskParams(
- model="llama-3.2-1b",
- messages=[
- ChatCompletionMessage(role="user", content="Hello, how are you?")
- ],
+ await command_sender.send(
+ ForwarderCommand(
+ origin=node_id,
+ tagged_command=TaggedCommand.from_(
+ ChatCompletion(
+ command_id=CommandId(),
+ request_params=ChatCompletionTaskParams(
+ model="llama-3.2-1b",
+ messages=[
+ ChatCompletionMessage(
+ role="user", content="Hello, how are you?"
+ )
+ ],
+ ),
+ )
),
)
)
- while len(await _get_events()) < 4:
+ while len(await _get_events()) < 3:
await asyncio.sleep(0.001)
events = await _get_events()
- print(events)
- assert len(events) == 4
- assert events[0].idx == 1
- assert isinstance(events[0].event, TopologyNodeCreated)
- assert isinstance(events[1].event, NodePerformanceMeasured)
- assert isinstance(events[2].event, InstanceCreated)
- runner_id = list(events[2].event.instance.shard_assignments.runner_to_shard.keys())[
+ assert len(events) == 3
+ assert events[0].idx == 0
+ assert events[1].idx == 1
+ assert events[2].idx == 2
+ assert isinstance(events[0].event, NodePerformanceMeasured)
+ assert isinstance(events[1].event, InstanceCreated)
+ runner_id = list(events[1].event.instance.shard_assignments.runner_to_shard.keys())[
0
]
- assert events[2].event == InstanceCreated(
+ assert events[1].event == InstanceCreated(
+ event_id=events[1].event.event_id,
instance=Instance(
- instance_id=events[2].event.instance.instance_id,
+ instance_id=events[1].event.instance.instance_id,
instance_type=InstanceStatus.ACTIVE,
shard_assignments=ShardAssignments(
- model_id="llama-3.2-1b",
+ model_id=ModelId("llama-3.2-1b"),
runner_to_shard={
(runner_id): PipelineShardMetadata(
partition_strategy=PartitionStrategy.pipeline,
@@ -164,10 +162,10 @@ async def test_master():
end_layer=16,
n_layers=16,
model_meta=ModelMetadata(
- model_id="llama-3.2-1b",
+ model_id=ModelId("llama-3.2-1b"),
pretty_name="Llama 3.2 1B",
n_layers=16,
- storage_size_kilobytes=678948,
+ storage_size=Memory.from_bytes(678948),
),
device_rank=0,
world_size=1,
@@ -176,16 +174,17 @@ async def test_master():
node_to_runner={node_id: runner_id},
),
hosts=[],
- )
+ ),
)
- assert isinstance(events[3].event, TaskCreated)
- assert events[3].event == TaskCreated(
- task_id=events[3].event.task_id,
+ assert isinstance(events[2].event, TaskCreated)
+ assert events[2].event == TaskCreated(
+ event_id=events[2].event.event_id,
+ task_id=events[2].event.task_id,
task=ChatCompletionTask(
- task_id=events[3].event.task_id,
- command_id=events[3].event.task.command_id,
+ task_id=events[2].event.task_id,
+ command_id=events[2].event.task.command_id,
task_type=TaskType.CHAT_COMPLETION,
- instance_id=events[3].event.task.instance_id,
+ instance_id=events[2].event.task.instance_id,
task_status=TaskStatus.PENDING,
task_params=ChatCompletionTaskParams(
model="llama-3.2-1b",
@@ -195,4 +194,3 @@ async def test_master():
),
),
)
- assert len(command_buffer) == 0
diff --git a/src/exo/master/tests/test_placement.py b/src/exo/master/tests/test_placement.py
index 16a33200..6b3aabf6 100644
--- a/src/exo/master/tests/test_placement.py
+++ b/src/exo/master/tests/test_placement.py
@@ -66,7 +66,7 @@ def test_get_instance_placements_create_instance(
expected_layers: tuple[int, int, int],
topology: Topology,
model_meta: ModelMetadata,
- create_node: Callable[[Memory, NodeId | None], NodeInfo],
+ create_node: Callable[[int, NodeId | None], NodeInfo],
create_connection: Callable[[NodeId, NodeId], Connection],
):
# arrange
@@ -82,9 +82,9 @@ def test_get_instance_placements_create_instance(
node_id_a = NodeId()
node_id_b = NodeId()
node_id_c = NodeId()
- topology.add_node(create_node(Memory.from_bytes(available_memory[0]), node_id_a))
- topology.add_node(create_node(Memory.from_bytes(available_memory[1]), node_id_b))
- topology.add_node(create_node(Memory.from_bytes(available_memory[2]), node_id_c))
+ topology.add_node(create_node(available_memory[0], node_id_a))
+ topology.add_node(create_node(available_memory[1], node_id_b))
+ topology.add_node(create_node(available_memory[2], node_id_c))
topology.add_connection(create_connection(node_id_a, node_id_b))
topology.add_connection(create_connection(node_id_b, node_id_c))
topology.add_connection(create_connection(node_id_c, node_id_a))
diff --git a/src/exo/master/tests/test_placement_utils.py b/src/exo/master/tests/test_placement_utils.py
index 31796a36..3b177a0e 100644
--- a/src/exo/master/tests/test_placement_utils.py
+++ b/src/exo/master/tests/test_placement_utils.py
@@ -252,9 +252,9 @@ def test_get_hosts_from_subgraph(
# assert
assert len(hosts) == 3
expected_hosts = [
- Host(ip=("127.0.0.1"), port=5001),
- Host(ip=("127.0.0.1"), port=5002),
- Host(ip=("127.0.0.1"), port=5003),
+ Host(ip=("169.254.0.1"), port=5001),
+ Host(ip=("169.254.0.1"), port=5002),
+ Host(ip=("169.254.0.1"), port=5003),
]
for expected_host in expected_hosts:
assert expected_host in hosts
diff --git a/src/exo/utils/channels.py b/src/exo/utils/channels.py
index bc203e53..b7a68bff 100644
--- a/src/exo/utils/channels.py
+++ b/src/exo/utils/channels.py
@@ -14,7 +14,7 @@ from anyio.streams.memory import (
class Sender[T](AnyioSender[T]):
def clone_receiver(self) -> "Receiver[T]":
- """Constructs a Sender using a Receivers shared state - similar to calling Receiver.clone() without needing the receiver"""
+ """Constructs a Receiver using a Senders shared state - similar to calling Receiver.clone() without needing the receiver"""
if self._closed:
raise ClosedResourceError
return Receiver(_state=self._state)
← 962e5ef4 version bump for brew consistency
·
back to Exo
·
Disable build macos app e01f9cf7 →