[object Object]

← 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

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 →