← back to Exo
Fixes for running this end to end
57ca487fdefb7b7b7715d637f847a3b30354c396 · 2025-07-28 10:51:03 +0100 · Alex Cheema
Co-authored-by: Gelu Vrabie <gelu.vrabie.univ@gmail.com>
Co-authored-by: Gelu Vrabie <gelu@exolabs.net>
Files touched
M master/api.pyM master/forwarder_supervisor.pyM master/main.pyM master/placement.pyM master/tests/test_forwarder_manager.pyM master/tests/test_master.pyM master/tests/test_placement.pyM master/tests/test_placement_utils.pyM rust/discovery/Cargo.tomlM rust/discovery/src/lib.rsM rust/exo_pyo3_bindings/src/discovery.rsM shared/apply/apply.pyM shared/topology.pyM shared/types/tasks.pyM worker/download/download_utils.pyM worker/download/impl_shard_downloader.pyM worker/download/shard_downloader.pyM worker/main.pyM worker/runner/runner_supervisor.pyM worker/tests/conftest.pyM worker/tests/test_worker_integration.pyM worker/tests/test_worker_plan.pyM worker/tests/test_worker_plan_utils.py
Diff
commit 57ca487fdefb7b7b7715d637f847a3b30354c396
Author: Alex Cheema <41707476+AlexCheema@users.noreply.github.com>
Date: Mon Jul 28 10:51:03 2025 +0100
Fixes for running this end to end
Co-authored-by: Gelu Vrabie <gelu.vrabie.univ@gmail.com>
Co-authored-by: Gelu Vrabie <gelu@exolabs.net>
---
master/api.py | 56 ++++++++++-------
master/forwarder_supervisor.py | 6 ++
master/main.py | 18 ++++--
master/placement.py | 6 +-
master/tests/test_forwarder_manager.py | 10 +--
master/tests/test_master.py | 101 ++++++++++++++++++++++++++++---
master/tests/test_placement.py | 14 +++--
master/tests/test_placement_utils.py | 26 ++++----
rust/discovery/Cargo.toml | 1 +
rust/discovery/src/lib.rs | 6 +-
rust/exo_pyo3_bindings/src/discovery.rs | 87 ++++++++++++++++++--------
shared/apply/apply.py | 5 ++
shared/topology.py | 9 ++-
shared/types/tasks.py | 3 +-
worker/download/download_utils.py | 26 +++++---
worker/download/impl_shard_downloader.py | 10 +++
worker/download/shard_downloader.py | 19 ++++++
worker/main.py | 98 +++++++++++++++++++++++-------
worker/runner/runner_supervisor.py | 2 +-
worker/tests/conftest.py | 10 ++-
worker/tests/test_worker_integration.py | 11 ++--
worker/tests/test_worker_plan.py | 34 +++++------
worker/tests/test_worker_plan_utils.py | 16 +----
23 files changed, 412 insertions(+), 162 deletions(-)
diff --git a/master/api.py b/master/api.py
index 387f2e5d..5b4ea986 100644
--- a/master/api.py
+++ b/master/api.py
@@ -29,6 +29,7 @@ from shared.types.events.commands import (
DeleteInstanceCommand,
)
from shared.types.events.components import EventFromEventLog
+from shared.types.models import ModelMetadata
from shared.types.state import State
from shared.types.tasks import ChatCompletionTaskParams
from shared.types.worker.common import InstanceId
@@ -37,9 +38,9 @@ from shared.types.worker.instances import Instance
def chunk_to_response(chunk: TokenChunk) -> ChatCompletionResponse:
return ChatCompletionResponse(
- id='abc',
+ id=chunk.command_id,
created=int(time.time()),
- model='idk',
+ model=chunk.model,
choices=[
StreamingChoiceResponse(
index=0,
@@ -52,6 +53,12 @@ def chunk_to_response(chunk: TokenChunk) -> ChatCompletionResponse:
]
)
+async def resolve_model_meta(model_id: str) -> ModelMetadata:
+ if model_id in MODEL_CARDS:
+ model_card = MODEL_CARDS[model_id]
+ return model_card.metadata
+ else:
+ return await get_model_meta(model_id)
@final
class API:
@@ -67,7 +74,7 @@ class API:
# self._app.get("/topology/control_plane")(self.get_control_plane_topology)
# self._app.get("/topology/data_plane")(self.get_data_plane_topology)
# self._app.get("/instances/list")(self.list_instances)
- self._app.post("/instances/create")(self.create_instance)
+ self._app.post("/instance")(self.create_instance)
self._app.get("/instance/{instance_id}")(self.get_instance)
self._app.delete("/instance/{instance_id}")(self.delete_instance)
# self._app.get("/model/{model_id}/metadata")(self.get_model_data)
@@ -92,11 +99,7 @@ class API:
# return {"message": "Hello, World!"}
async def create_instance(self, payload: CreateInstanceTaskParams) -> CreateInstanceResponse:
- if payload.model_id in MODEL_CARDS:
- model_card = MODEL_CARDS[payload.model_id]
- model_meta = model_card.metadata
- else:
- model_meta = await get_model_meta(payload.model_id)
+ model_meta = await resolve_model_meta(payload.model_id)
command = CreateInstanceCommand(
command_id=CommandId(),
@@ -139,22 +142,12 @@ class API:
# def get_instances_by_model(self, model_id: ModelId) -> list[Instance]: ...
- async def _generate_chat_stream(self, payload: ChatCompletionTaskParams) -> AsyncGenerator[str, None]:
+ async def _generate_chat_stream(self, command_id: CommandId) -> AsyncGenerator[str, None]:
"""Generate chat completion stream as JSON strings."""
+
events = await self.global_events.get_events_since(0)
prev_idx = await self.global_events.get_last_idx()
- # At the moment, we just create the task in the API.
- # In the future, a `Request` will be created here and they will be bundled into `Task` objects by the master.
- command_id=CommandId()
-
- request = ChatCompletionCommand(
- command_id=command_id,
- command_type=CommandType.CHAT_COMPLETION,
- request_params=payload,
- )
- self.command_buffer.append(request)
-
finished = False
while not finished:
await asyncio.sleep(0.01)
@@ -177,10 +170,29 @@ class API:
return
+ async def _trigger_notify_user_to_download_model(self, model_id: str) -> None:
+ print("TODO: we should send a notification to the user to download the model")
+
async def chat_completions(self, payload: ChatCompletionTaskParams) -> StreamingResponse:
"""Handle chat completions with proper streaming response."""
+ model_meta = await resolve_model_meta(payload.model)
+ payload.model = model_meta.model_id
+
+ for instance in self.get_state().instances.values():
+ if instance.shard_assignments.model_id == payload.model:
+ break
+ else:
+ await self._trigger_notify_user_to_download_model(payload.model)
+ raise HTTPException(status_code=404, detail=f"No instance found for model {payload.model}")
+
+ command = ChatCompletionCommand(
+ command_id=CommandId(),
+ command_type=CommandType.CHAT_COMPLETION,
+ request_params=payload,
+ )
+ self.command_buffer.append(command)
return StreamingResponse(
- self._generate_chat_stream(payload),
+ self._generate_chat_stream(command.command_id),
media_type="text/plain"
)
@@ -195,4 +207,4 @@ def start_fastapi_server(
):
api = API(command_buffer, global_events, get_state)
- uvicorn.run(api.app, host=host, port=port)
\ No newline at end of file
+ uvicorn.run(api.app, host=host, port=port)
diff --git a/master/forwarder_supervisor.py b/master/forwarder_supervisor.py
index 93a0bab0..d00f4418 100644
--- a/master/forwarder_supervisor.py
+++ b/master/forwarder_supervisor.py
@@ -10,6 +10,7 @@ from shared.constants import (
LIBP2P_GLOBAL_EVENTS_TOPIC,
LIBP2P_WORKER_EVENTS_TOPIC,
)
+from shared.types.common import NodeId
class ForwarderRole(str, Enum):
@@ -35,10 +36,12 @@ class ForwarderSupervisor:
def __init__(
self,
+ node_id: NodeId,
forwarder_binary_path: Path,
logger: Logger,
health_check_interval: float = 5.0
):
+ self.node_id = node_id
self._binary_path = forwarder_binary_path
self._logger = logger
self._health_check_interval = health_check_interval
@@ -108,6 +111,9 @@ class ForwarderSupervisor:
f'{pairs}',
stdout=None,
stderr=None,
+ env={
+ "FORWARDER_NODE_ID": str(self.node_id),
+ }
)
self._logger.info(f"Starting forwarder with forwarding pairs: {pairs}")
diff --git a/master/main.py b/master/main.py
index 24868af7..6417b9c4 100644
--- a/master/main.py
+++ b/master/main.py
@@ -39,6 +39,7 @@ class Master:
def __init__(self, node_id_keypair: Keypair, node_id: NodeId, command_buffer: list[Command],
global_events: AsyncSQLiteEventStorage, worker_events: AsyncSQLiteEventStorage,
forwarder_binary_path: Path, logger: logging.Logger):
+ self.state = State()
self.node_id = node_id
self.command_buffer = command_buffer
self.global_events = global_events
@@ -46,16 +47,22 @@ class Master:
self.discovery_supervisor = DiscoverySupervisor(
node_id_keypair,
node_id,
- global_events,
+ # TODO: needs to be more general for when we have master election
+ worker_events if os.getenv('EXO_RUN_AS_REPLICA') in set(['TRUE', 'true', '1']) else global_events,
logger
)
self.forwarder_supervisor = ForwarderSupervisor(
+ self.node_id,
forwarder_binary_path=forwarder_binary_path,
logger=logger
)
self.election_callbacks = ElectionCallbacks(self.forwarder_supervisor, logger)
self.logger = logger
+ @property
+ def event_log_for_reads(self) -> AsyncSQLiteEventStorage:
+ return self.global_events
+
@property
def event_log_for_writes(self) -> AsyncSQLiteEventStorage:
if self.forwarder_supervisor.current_role == ForwarderRole.MASTER:
@@ -89,8 +96,9 @@ class Master:
next_events.append(TaskCreated(
task_id=task_id,
task=ChatCompletionTask(
- task_id=task_id,
task_type=TaskType.CHAT_COMPLETION,
+ task_id=task_id,
+ command_id=next_command.command_id,
instance_id=matching_instance.instance_id,
task_status=TaskStatus.PENDING,
task_params=next_command.request_params
@@ -108,16 +116,18 @@ class Master:
await self.event_log_for_writes.append_events(next_events, origin=self.node_id)
# 2. get latest events
- events = await self.global_events.get_events_since(self.state.last_event_applied_idx)
+ events = await self.event_log_for_reads.get_events_since(self.state.last_event_applied_idx)
if len(events) == 0:
await asyncio.sleep(0.01)
return
+ self.logger.info(f"got events: {events}")
# 3. for each event, apply it to the state
for event_from_log in events:
+ print(f"applying event: {event_from_log}")
self.state = apply(self.state, event_from_log)
- self.logger.info(f"state: {self.state}")
+ self.logger.info(f"state: {self.state.model_dump_json()}")
async def run(self):
self.state = await self._get_state_snapshot()
diff --git a/master/placement.py b/master/placement.py
index 7137938f..cd3320cc 100644
--- a/master/placement.py
+++ b/master/placement.py
@@ -28,10 +28,10 @@ def get_instance_placements(
all_nodes = list(topology.list_nodes())
cycles = topology.get_cycles()
- nodes_in_cycles = {node.node_id for cycle in cycles for node in cycle}
- singleton_cycles = [[node] for node in all_nodes if node.node_id not in nodes_in_cycles]
+ # we can also always just have a node on its own
+ singleton_cycles = [[node] for node in all_nodes]
candidate_cycles = cycles + singleton_cycles
- cycles_with_sufficient_memory = filter_cycles_by_memory(candidate_cycles, command.model_meta.storage_size_kilobytes)
+ cycles_with_sufficient_memory = filter_cycles_by_memory(candidate_cycles, command.model_meta.storage_size_kilobytes * 1024)
if not cycles_with_sufficient_memory:
raise ValueError("No cycles found with sufficient memory")
diff --git a/master/tests/test_forwarder_manager.py b/master/tests/test_forwarder_manager.py
index 0160362b..c9413c52 100644
--- a/master/tests/test_forwarder_manager.py
+++ b/master/tests/test_forwarder_manager.py
@@ -24,6 +24,7 @@ from shared.constants import (
LIBP2P_GLOBAL_EVENTS_TOPIC,
LIBP2P_WORKER_EVENTS_TOPIC,
)
+from shared.types.common import NodeId
# Mock forwarder script content
MOCK_FORWARDER_SCRIPT = '''#!/usr/bin/env python3
@@ -182,7 +183,7 @@ class TestForwardersupervisorBasic:
# Set environment
os.environ.update(mock_env_vars)
- supervisor = ForwarderSupervisor(mock_forwarder_script, test_logger)
+ supervisor = ForwarderSupervisor(NodeId(), mock_forwarder_script, test_logger)
await supervisor.start_as_replica()
# Track the process for cleanup
@@ -224,7 +225,7 @@ class TestForwardersupervisorBasic:
"""Test changing role from replica to master."""
os.environ.update(mock_env_vars)
- supervisor = ForwarderSupervisor(mock_forwarder_script, test_logger)
+ supervisor = ForwarderSupervisor(NodeId(), mock_forwarder_script, test_logger)
await supervisor.start_as_replica()
if supervisor.process:
@@ -268,7 +269,7 @@ class TestForwardersupervisorBasic:
"""Test that setting the same role twice doesn't restart the process."""
os.environ.update(mock_env_vars)
- supervisor = ForwarderSupervisor(mock_forwarder_script, test_logger)
+ supervisor = ForwarderSupervisor(NodeId(), mock_forwarder_script, test_logger)
await supervisor.start_as_replica()
original_pid = supervisor.process_pid
@@ -300,6 +301,7 @@ class TestForwardersupervisorBasic:
os.environ.update(mock_env_vars)
supervisor = ForwarderSupervisor(
+ NodeId(),
mock_forwarder_script,
test_logger,
health_check_interval=0.5 # Faster health checks for testing
@@ -346,7 +348,7 @@ class TestForwardersupervisorBasic:
"""Test behavior when forwarder binary doesn't exist."""
nonexistent_path = temp_dir / "nonexistent_forwarder"
- supervisor = ForwarderSupervisor(nonexistent_path, test_logger)
+ supervisor = ForwarderSupervisor(NodeId(), nonexistent_path, test_logger)
# Should raise FileNotFoundError
with pytest.raises(FileNotFoundError):
diff --git a/master/tests/test_master.py b/master/tests/test_master.py
index 4c4d23e4..767481e9 100644
--- a/master/tests/test_master.py
+++ b/master/tests/test_master.py
@@ -14,9 +14,27 @@ from shared.db.sqlite.event_log_manager import EventLogManager
from shared.types.api import ChatCompletionMessage, ChatCompletionTaskParams
from shared.types.common import NodeId
from shared.types.events import TaskCreated
-from shared.types.events._events import TopologyNodeCreated
-from shared.types.events.commands import ChatCompletionCommand, Command, CommandId
+from shared.types.events._events import (
+ InstanceCreated,
+ NodePerformanceMeasured,
+ TopologyNodeCreated,
+)
+from shared.types.events.commands import (
+ ChatCompletionCommand,
+ Command,
+ CommandId,
+ CreateInstanceCommand,
+)
+from shared.types.models import ModelMetadata
+from shared.types.profiling import (
+ MemoryPerformanceProfile,
+ NodePerformanceProfile,
+ SystemPerformanceProfile,
+)
from shared.types.tasks import ChatCompletionTask, TaskStatus, TaskType
+from shared.types.worker.common import InstanceId
+from shared.types.worker.instances import Instance, InstanceStatus, ShardAssignments
+from shared.types.worker.shards import PartitionStrategy, PipelineShardMetadata
def _create_forwarder_dummy_binary() -> Path:
@@ -42,9 +60,40 @@ async def test_master():
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, logger=logger, worker_events=event_log_manager.worker_events)
+ forwarder_binary_path=forwarder_binary_path, logger=logger, worker_events=global_events)
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)
+ # 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",
+ memory=MemoryPerformanceProfile(ram_total=678948*1024, ram_available=678948*1024, swap_total=0, swap_available=0),
+ network_interfaces=[],
+ system=SystemPerformanceProfile(flops_fp16=0)
+ )
+ )
+ ], origin=node_id)
+ while len(master.state.node_profiles) == 0:
+ await asyncio.sleep(0.001)
+ command_buffer.append(CreateInstanceCommand(
+ command_id=CommandId(),
+ instance_id=InstanceId(),
+ model_meta=ModelMetadata(
+ model_id="llama-3.2-1b",
+ pretty_name="Llama 3.2 1B",
+ n_layers=16,
+ storage_size_kilobytes=678948
+ )
+ ))
+ while len(master.state.instances.keys()) == 0:
+ await asyncio.sleep(0.001)
command_buffer.append(
ChatCompletionCommand(
command_id=CommandId(),
@@ -54,20 +103,52 @@ async def test_master():
)
)
)
- while len(await global_events.get_events_since(0)) == 0:
+ while len(await global_events.get_events_since(0)) < 4:
await asyncio.sleep(0.001)
events = await global_events.get_events_since(0)
- assert len(events) == 2
+ print(events)
+ assert len(events) == 4
assert events[0].idx_in_log == 1
assert isinstance(events[0].event, TopologyNodeCreated)
- assert isinstance(events[1].event, TaskCreated)
- assert events[1].event == TaskCreated(
- task_id=events[1].event.task_id,
+ 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())[0]
+ assert events[2].event == InstanceCreated(
+ instance=Instance(
+ instance_id=events[2].event.instance.instance_id,
+ instance_type=InstanceStatus.ACTIVE,
+ shard_assignments=ShardAssignments(
+ model_id="llama-3.2-1b",
+ runner_to_shard={
+ (runner_id): PipelineShardMetadata(
+ partition_strategy=PartitionStrategy.pipeline,
+ start_layer=0,
+ end_layer=16,
+ n_layers=16,
+ model_meta=ModelMetadata(
+ model_id="llama-3.2-1b",
+ pretty_name="Llama 3.2 1B",
+ n_layers=16,
+ storage_size_kilobytes=678948
+ ),
+ device_rank=0,
+ world_size=1
+ )
+ },
+ 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,
task=ChatCompletionTask(
- task_id=events[1].event.task_id,
+ task_id=events[3].event.task_id,
+ command_id=events[3].event.task.command_id,
task_type=TaskType.CHAT_COMPLETION,
- instance_id=events[1].event.task.instance_id,
+ instance_id=events[3].event.task.instance_id,
task_status=TaskStatus.PENDING,
task_params=ChatCompletionTaskParams(
model="llama-3.2-1b",
diff --git a/master/tests/test_placement.py b/master/tests/test_placement.py
index d51d16b1..1ab8a9ef 100644
--- a/master/tests/test_placement.py
+++ b/master/tests/test_placement.py
@@ -64,6 +64,8 @@ def test_get_instance_placements_create_instance(
create_node: Callable[[int, NodeId | None], Node],
create_connection: Callable[[NodeId, NodeId], Connection]
):
+ # TODO: this test is not exactly what we want. if a model can fit on one node, it should be placed there.
+ # TODO: right now we assume it will be placed across all nodes.
# arrange
model_meta.n_layers = total_layers
@@ -75,9 +77,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(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_node(create_node(available_memory[0]*1024, node_id_a))
+ topology.add_node(create_node(available_memory[1]*1024, node_id_b))
+ topology.add_node(create_node(available_memory[2]*1024, 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))
@@ -113,7 +115,7 @@ def test_get_instance_placements_one_node_exact_fit(
) -> None:
topology = Topology()
node_id = NodeId()
- topology.add_node(create_node(1000, node_id))
+ topology.add_node(create_node(1000*1024, node_id))
create_instance_command = CreateInstanceCommand(
command_id=CommandId(),
model_meta=ModelMetadata(
@@ -139,7 +141,7 @@ def test_get_instance_placements_one_node_fits_with_extra_memory(
) -> None:
topology = Topology()
node_id = NodeId()
- topology.add_node(create_node(1001, node_id))
+ topology.add_node(create_node(1001*1024, node_id))
create_instance_command = CreateInstanceCommand(
command_id=CommandId(),
model_meta=ModelMetadata(
@@ -165,7 +167,7 @@ def test_get_instance_placements_one_node_not_fit(
) -> None:
topology = Topology()
node_id = NodeId()
- topology.add_node(create_node(1000, node_id))
+ topology.add_node(create_node(1000*1024, node_id))
create_instance_command = CreateInstanceCommand(
command_id=CommandId(),
model_meta=ModelMetadata(
diff --git a/master/tests/test_placement_utils.py b/master/tests/test_placement_utils.py
index 2ef84cd1..f9c286d9 100644
--- a/master/tests/test_placement_utils.py
+++ b/master/tests/test_placement_utils.py
@@ -24,8 +24,8 @@ def test_filter_cycles_by_memory(topology: Topology, create_node: Callable[[int,
node1_id = NodeId()
node2_id = NodeId()
- node1 = create_node(1000, node1_id)
- node2 = create_node(1000, node2_id)
+ node1 = create_node(1000*1024, node1_id)
+ node2 = create_node(1000*1024, node2_id)
topology.add_node(node1)
topology.add_node(node2)
@@ -52,8 +52,8 @@ def test_filter_cycles_by_insufficient_memory(topology: Topology, create_node: C
node1_id = NodeId()
node2_id = NodeId()
- node1 = create_node(1000, node1_id)
- node2 = create_node(1000, node2_id)
+ node1 = create_node(1000*1024, node1_id)
+ node2 = create_node(1000*1024, node2_id)
topology.add_node(node1)
topology.add_node(node2)
@@ -77,9 +77,9 @@ def test_filter_multiple_cycles_by_memory(topology: Topology, create_node: Calla
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)
+ node_a = create_node(500*1024, node_a_id)
+ node_b = create_node(500*1024, node_b_id)
+ node_c = create_node(1000*1024, node_c_id)
topology.add_node(node_a)
topology.add_node(node_b)
@@ -107,9 +107,9 @@ def test_get_smallest_cycles(topology: Topology, create_node: Callable[[int, Nod
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)
+ node_a = create_node(500*1024, node_a_id)
+ node_b = create_node(500*1024, node_b_id)
+ node_c = create_node(1000*1024, node_c_id)
topology.add_node(node_a)
topology.add_node(node_b)
@@ -139,9 +139,9 @@ def test_get_shard_assignments(topology: Topology, create_node: Callable[[int, N
node_b_id = NodeId()
node_c_id = NodeId()
- node_a = create_node(available_memory[0], node_a_id)
- node_b = create_node(available_memory[1], node_b_id)
- node_c = create_node(available_memory[2], node_c_id)
+ node_a = create_node(available_memory[0]*1024, node_a_id)
+ node_b = create_node(available_memory[1]*1024, node_b_id)
+ node_c = create_node(available_memory[2]*1024, node_c_id)
topology.add_node(node_a)
topology.add_node(node_b)
diff --git a/rust/discovery/Cargo.toml b/rust/discovery/Cargo.toml
index 6ca9ef17..ff94a8be 100644
--- a/rust/discovery/Cargo.toml
+++ b/rust/discovery/Cargo.toml
@@ -33,6 +33,7 @@ thiserror = { workspace = true }
#itertools = { workspace = true }
tracing-subscriber = { version = "0.3.19", features = ["default", "env-filter"] }
keccak-const = { workspace = true }
+log = "0.4"
# Networking
libp2p = { workspace = true, features = ["full"] }
\ No newline at end of file
diff --git a/rust/discovery/src/lib.rs b/rust/discovery/src/lib.rs
index 17cb78ca..bcc1075a 100644
--- a/rust/discovery/src/lib.rs
+++ b/rust/discovery/src/lib.rs
@@ -41,6 +41,8 @@ pub(crate) mod private {
/// Create and configure a swarm, and start listening to all ports/OS.
#[inline]
pub fn discovery_swarm(keypair: identity::Keypair) -> alias::AnyResult<Swarm<DiscoveryBehaviour>> {
+ let peer_id = keypair.public().to_peer_id();
+ log::info!("RUST: Creating discovery swarm with peer_id: {}", peer_id);
let mut swarm = SwarmBuilder::with_existing_identity(keypair)
.with_tokio()
.with_other_transport(discovery_transport)?
@@ -49,7 +51,9 @@ pub fn discovery_swarm(keypair: identity::Keypair) -> alias::AnyResult<Swarm<Dis
// Listen on all interfaces and whatever port the OS assigns
// swarm.listen_on("/ip4/0.0.0.0/udp/0/quic-v1".parse()?)?; // TODO: make this
- swarm.listen_on("/ip4/0.0.0.0/tcp/0".parse()?)?;
+ let listen_addr = "/ip4/0.0.0.0/tcp/0".parse()?;
+ log::info!("RUST: Attempting to listen on: {}", listen_addr);
+ swarm.listen_on(listen_addr)?;
Ok(swarm)
}
diff --git a/rust/exo_pyo3_bindings/src/discovery.rs b/rust/exo_pyo3_bindings/src/discovery.rs
index 411c41b6..3ba8bbc6 100644
--- a/rust/exo_pyo3_bindings/src/discovery.rs
+++ b/rust/exo_pyo3_bindings/src/discovery.rs
@@ -120,11 +120,29 @@ async fn discovery_task(
Behaviour(Mdns(Discovered(list))) => {
for (peer_id, multiaddr) in list {
log::info!("RUST: mDNS discovered a new peer: {peer_id} on {multiaddr}");
- // TODO: this does the job of (actually) creating & maintaining connection
- // but its coupled to gossipsub & also the connection isn't configured
- // for setting "connection keep alive" in NetworkBehavior's ConnectionHandler
- // >in future, make own small NetworkBehavior impl just to track this state
+ let local_peer_id = *swarm.local_peer_id();
+ // To avoid simultaneous dial races, only the lexicographically larger peer_id dials.
+ if peer_id > local_peer_id {
+ let dial_opts = DialOpts::peer_id(peer_id)
+ .addresses(vec![multiaddr.clone()].into())
+ .condition(libp2p::swarm::dial_opts::PeerCondition::Always)
+ .build();
+ match swarm.dial(dial_opts) {
+ Ok(()) => log::info!("RUST: Dial initiated to {multiaddr}"),
+ Err(libp2p::swarm::DialError::DialPeerConditionFalse(_)) => {
+ // Another dial is already in progress; not an error for us.
+ log::debug!(
+ "RUST: Dial skipped because another dial is active for {peer_id}"
+ );
+ }
+ Err(e) => {
+ log::warn!("RUST: Failed to dial {multiaddr}: {e:?}");
+ }
+ }
+ }
+ // Maintain peer in gossipsub mesh so the connection stays alive once established.
swarm.behaviour_mut().gossipsub.add_explicit_peer(&peer_id);
+ log::info!("RUST: Added peer {peer_id} to gossipsub explicit peers");
}
}
Behaviour(Mdns(Expired(list))) => {
@@ -149,6 +167,7 @@ async fn discovery_task(
concurrent_dial_errors,
established_in: _established_in,
} => {
+ log::info!("RUST: ConnectionEstablished event - peer_id: {peer_id}, connection_id: {connection_id:?}, endpoint: {endpoint:?}");
// log any connection errors
if let Some(concurrent_dial_errors) = concurrent_dial_errors {
for (multiaddr, error) in concurrent_dial_errors {
@@ -156,17 +175,21 @@ async fn discovery_task(
}
}
- // TODO: right now we assume we are using TCP/IP which treats all nodes
- // as both dialers AND listeners. This means for each connection you will actually
- // see TWO duplicate Connected events => Dialer & Listener
- // SO ignore the Dialer & extract the info we need from Listener
- // HOWEVER this makes the swarm implicitly rely on TCP/IP, so is brittle to changes
- // e.g. adding QUIC protocol or something
- // >As soon as we add anything other than TCP/IP, this must be updated or there will be broken code
- let ConnectedPoint::Listener { local_addr, send_back_addr } = endpoint else {
- log::warn!("Ignoring `ConnectedPoint::Dialer` event because for TCP/IP it has a dual `ConnectedPoint::Listener` event: {endpoint:?}");
- continue;
+ // Extract addresses based on endpoint type
+ let (local_addr, send_back_addr) = match &endpoint {
+ ConnectedPoint::Listener { local_addr, send_back_addr } => {
+ log::info!("RUST: Connection established (Listener) - local_addr: {local_addr}, send_back_addr: {send_back_addr}");
+ (local_addr.clone(), send_back_addr.clone())
+ },
+ ConnectedPoint::Dialer { address, .. } => {
+ log::info!("RUST: Connection established (Dialer) - remote_addr: {address}");
+ // For dialer, we use the dialed address as both local and send_back
+ // This isn't perfect but allows both sides to be notified
+ (address.clone(), address.clone())
+ }
};
+
+ log::info!("RUST: Number of connected callbacks: {}", connected_callbacks.len());
// trigger callback on connected peer
@@ -180,22 +203,27 @@ async fn discovery_task(
}
},
ConnectionClosed { peer_id, connection_id, endpoint, num_established, cause } => {
+ log::info!("RUST: ConnectionClosed event - peer_id: {peer_id}, connection_id: {connection_id:?}, endpoint: {endpoint:?}, num_established: {num_established}");
// log any connection errors
if let Some(cause) = cause {
log::error!("Connection error: cause={cause:?}");
}
- // TODO: right now we assume we are using TCP/IP which treats all nodes
- // as both dialers AND listeners. This means for each connection you will actually
- // see TWO duplicate Connected events => Dialer & Listener
- // SO ignore the Dialer & extract the info we need from Listener
- // HOWEVER this makes the swarm implicitly rely on TCP/IP, so is brittle to changes
- // e.g. adding QUIC protocol or something
- // >As soon as we add anything other than TCP/IP, this must be updated or there will be broken code
- let ConnectedPoint::Listener { local_addr, send_back_addr } = endpoint else {
- log::warn!("Ignoring `ConnectedPoint::Dialer` event because for TCP/IP it has a dual `ConnectedPoint::Listener` event: {endpoint:?}");
- continue;
+ // Extract addresses based on endpoint type
+ let (local_addr, send_back_addr) = match &endpoint {
+ ConnectedPoint::Listener { local_addr, send_back_addr } => {
+ log::info!("RUST: Connection closed (Listener) - local_addr: {local_addr}, send_back_addr: {send_back_addr}");
+ (local_addr.clone(), send_back_addr.clone())
+ },
+ ConnectedPoint::Dialer { address, .. } => {
+ log::info!("RUST: Connection closed (Dialer) - remote_addr: {address}");
+ // For dialer, we use the dialed address as both local and send_back
+ // This isn't perfect but allows both sides to be notified
+ (address.clone(), address.clone())
+ }
};
+
+ log::info!("RUST: Number of disconnected callbacks: {}", disconnected_callbacks.len());
// trigger callback on connected peer
for disconnected_callback in &disconnected_callbacks {
@@ -207,8 +235,13 @@ async fn discovery_task(
});
}
}
+ NewListenAddr { address, .. } => {
+ log::info!("RUST: Local node is listening on {address}");
+ let local_peer = swarm.local_peer_id();
+ log::info!("RUST: Local peer_id: {local_peer}");
+ }
e => {
- log::info!("RUST: Other event {e:?}");
+ log::debug!("RUST: Other event {e:?}");
}
}
}
@@ -258,15 +291,19 @@ impl PyDiscoveryService {
// get identity
let identity = identity.borrow().0.clone();
+ log::info!("RUST: Creating DiscoveryService with keypair");
// create discovery swarm (within tokio context!! or it crashes)
let swarm = get_runtime()
.block_on(async { discovery_swarm(identity) })
.pyerr()?;
+ log::info!("RUST: Discovery swarm created successfully");
// spawn tokio task
get_runtime().spawn(async move {
+ log::info!("RUST: Starting discovery task");
discovery_task(receiver, swarm).await;
+ log::info!("RUST: Discovery task ended");
});
Ok(Self::new(sender))
}
diff --git a/shared/apply/apply.py b/shared/apply/apply.py
index b5f49538..1386a475 100644
--- a/shared/apply/apply.py
+++ b/shared/apply/apply.py
@@ -114,6 +114,9 @@ def apply_runner_deleted(event: RunnerDeleted, state: State) -> State:
def apply_node_performance_measured(event: NodePerformanceMeasured, state: State) -> State:
new_profiles: Mapping[NodeId, NodePerformanceProfile] = {**state.node_profiles, event.node_id: event.node_profile}
state = state.model_copy(update={"node_profiles": new_profiles})
+ if not state.topology.contains_node(event.node_id):
+ # TODO: figure out why this is happening in the first place
+ return state
topology = copy.copy(state.topology)
topology.update_node_profile(event.node_id, event.node_profile)
return state.model_copy(update={"topology": topology})
@@ -148,5 +151,7 @@ def apply_topology_edge_replaced_atomically(event: TopologyEdgeReplacedAtomicall
@event_apply.register(TopologyEdgeDeleted)
def apply_topology_edge_deleted(event: TopologyEdgeDeleted, state: State) -> State:
topology = copy.copy(state.topology)
+ if not topology.contains_connection(event.edge):
+ return state
topology.remove_connection(event.edge)
return state.model_copy(update={"topology": topology})
\ No newline at end of file
diff --git a/shared/topology.py b/shared/topology.py
index d007e532..2263c447 100644
--- a/shared/topology.py
+++ b/shared/topology.py
@@ -63,6 +63,12 @@ class Topology(TopologyProto):
self._node_id_to_rx_id_map[node.node_id] = rx_id
self._rx_id_to_node_id_map[rx_id] = node.node_id
+ def contains_node(self, node_id: NodeId) -> bool:
+ return node_id in self._node_id_to_rx_id_map
+
+ def contains_connection(self, connection: Connection) -> bool:
+ return connection in self._edge_id_to_rx_id_map
+
def add_connection(
self,
connection: Connection,
@@ -120,7 +126,8 @@ class Topology(TopologyProto):
else:
self._graph.remove_edge_from_index(rx_idx)
del self._edge_id_to_rx_id_map[connection]
- del self._rx_id_to_node_id_map[rx_idx]
+ if rx_idx in self._rx_id_to_node_id_map:
+ del self._rx_id_to_node_id_map[rx_idx]
def get_cycles(self) -> list[list[Node]]:
cycle_idxs = rx.simple_cycles(self._graph)
diff --git a/shared/types/tasks.py b/shared/types/tasks.py
index 12b0b514..00426ba9 100644
--- a/shared/types/tasks.py
+++ b/shared/types/tasks.py
@@ -4,7 +4,7 @@ from typing import Annotated, Literal
from pydantic import BaseModel, Field
from shared.types.api import ChatCompletionTaskParams
-from shared.types.common import ID
+from shared.types.common import ID, CommandId
from shared.types.worker.common import InstanceId
@@ -26,6 +26,7 @@ class TaskStatus(str, Enum):
class ChatCompletionTask(BaseModel):
task_type: Literal[TaskType.CHAT_COMPLETION] = TaskType.CHAT_COMPLETION
task_id: TaskId
+ command_id: CommandId
instance_id: InstanceId
task_status: TaskStatus
task_params: ChatCompletionTaskParams
diff --git a/worker/download/download_utils.py b/worker/download/download_utils.py
index cde8f056..a5615163 100644
--- a/worker/download/download_utils.py
+++ b/worker/download/download_utils.py
@@ -198,17 +198,22 @@ async def calc_hash(path: Path, hash_type: Literal["sha1", "sha256"] = "sha1") -
hasher.update(chunk)
return hasher.hexdigest()
-async def file_meta(repo_id: str, revision: str, path: str) -> Tuple[int, str]:
- url = urljoin(f"{get_hf_endpoint()}/{repo_id}/resolve/{revision}/", path)
+async def file_meta(repo_id: str, revision: str, path: str, redirected_location: str | None = None) -> Tuple[int, str]:
+ # NOTE: huggingface broke the E-Tag so we can no longer assume E-Tag == sha256(file)
+ url = urljoin(f"{get_hf_endpoint()}/{repo_id}/resolve/{revision}/", path) if redirected_location is None else f"{get_hf_endpoint()}{redirected_location}"
headers = await get_auth_headers()
async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=1800, connect=60, sock_read=1800, sock_connect=60)) as session, session.head(url, headers=headers) as r:
- content_length = int(r.headers.get('x-linked-size') or r.headers.get('content-length') or 0)
- etag = r.headers.get('X-Linked-ETag') or r.headers.get('ETag') or r.headers.get('Etag')
- assert content_length > 0, f"No content length for {url}"
- assert etag is not None, f"No remote hash for {url}"
- if (etag[0] == '"' and etag[-1] == '"') or (etag[0] == "'" and etag[-1] == "'"):
- etag = etag[1:-1]
- return content_length, etag
+ if r.status == 307:
+ redirected_location = r.headers.get('Location')
+ return await file_meta(repo_id, revision, path, redirected_location)
+
+ content_length = int(r.headers.get('x-linked-size') or r.headers.get('content-length') or 0)
+ etag = r.headers.get('X-Linked-ETag') or r.headers.get('ETag') or r.headers.get('Etag')
+ assert content_length > 0, f"No content length for {url}"
+ assert etag is not None, f"No remote hash for {url}"
+ if (etag[0] == '"' and etag[-1] == '"') or (etag[0] == "'" and etag[-1] == "'"):
+ etag = etag[1:-1]
+ return content_length, etag
async def download_file_with_retry(repo_id: str, revision: str, path: str, target_dir: Path, on_progress: Callable[[int, int], None] = lambda _, __: None) -> Path:
n_attempts = 30
@@ -243,7 +248,8 @@ async def _download_file(repo_id: str, revision: str, path: str, target_dir: Pat
assert r.status in [200, 206], f"Failed to download {path} from {url}: {r.status}"
async with aiofiles.open(partial_path, 'ab' if resume_byte_pos else 'wb') as f:
while chunk := await r.content.read(8 * 1024 * 1024):
- on_progress(n_read := n_read + await f.write(chunk), length)
+ n_read = n_read + (await f.write(chunk))
+ on_progress(n_read, length)
final_hash = await calc_hash(partial_path, hash_type="sha256" if len(remote_hash) == 64 else "sha1")
integrity = final_hash == remote_hash
diff --git a/worker/download/impl_shard_downloader.py b/worker/download/impl_shard_downloader.py
index 3843107e..dff56912 100644
--- a/worker/download/impl_shard_downloader.py
+++ b/worker/download/impl_shard_downloader.py
@@ -64,6 +64,9 @@ class SingletonShardDownloader(ShardDownloader):
async for path, status in self.shard_downloader.get_shard_download_status():
yield path, status
+ async def get_shard_download_status_for_shard(self, shard: ShardMetadata) -> RepoDownloadProgress:
+ return await self.shard_downloader.get_shard_download_status_for_shard(shard)
+
class CachedShardDownloader(ShardDownloader):
def __init__(self, shard_downloader: ShardDownloader):
self.shard_downloader = shard_downloader
@@ -86,6 +89,9 @@ class CachedShardDownloader(ShardDownloader):
async for path, status in self.shard_downloader.get_shard_download_status():
yield path, status
+ async def get_shard_download_status_for_shard(self, shard: ShardMetadata) -> RepoDownloadProgress:
+ return await self.shard_downloader.get_shard_download_status_for_shard(shard)
+
class ResumableShardDownloader(ShardDownloader):
def __init__(self, max_parallel_downloads: int = 8):
self.max_parallel_downloads = max_parallel_downloads
@@ -126,3 +132,7 @@ class ResumableShardDownloader(ShardDownloader):
yield (path, progress)
except Exception as e:
print("Error downloading shard:", e)
+
+ async def get_shard_download_status_for_shard(self, shard: ShardMetadata) -> RepoDownloadProgress:
+ _, progress = await download_shard(shard, self.on_progress_wrapper, skip_download=True)
+ return progress
diff --git a/worker/download/shard_downloader.py b/worker/download/shard_downloader.py
index 68a095c7..27b88411 100644
--- a/worker/download/shard_downloader.py
+++ b/worker/download/shard_downloader.py
@@ -68,6 +68,10 @@ class ShardDownloader(ABC):
)
)
+ @abstractmethod
+ async def get_shard_download_status_for_shard(self, shard: ShardMetadata) -> RepoDownloadProgress:
+ ...
+
class NoopShardDownloader(ShardDownloader):
async def ensure_shard(self, shard: ShardMetadata, config_only: bool = False) -> Path:
@@ -106,3 +110,18 @@ class NoopShardDownloader(ShardDownloader):
status="complete",
)
)
+
+ async def get_shard_download_status_for_shard(self, shard: ShardMetadata) -> RepoDownloadProgress:
+ return RepoDownloadProgress(
+ repo_id="noop",
+ repo_revision="noop",
+ shard=shard,
+ completed_files=0,
+ total_files=0,
+ downloaded_bytes=0,
+ downloaded_bytes_this_session=0,
+ total_bytes=0,
+ overall_speed=0,
+ overall_eta=timedelta(seconds=0),
+ status="complete",
+ )
\ No newline at end of file
diff --git a/worker/main.py b/worker/main.py
index 987da047..1275a3e6 100644
--- a/worker/main.py
+++ b/worker/main.py
@@ -1,8 +1,9 @@
import asyncio
import logging
-import os
from asyncio import Queue
+from copy import deepcopy
from functools import partial
+from time import process_time
from typing import AsyncGenerator, Optional
from pydantic import BaseModel, ConfigDict
@@ -54,7 +55,8 @@ from shared.types.worker.runners import (
)
from shared.types.worker.shards import ShardMetadata
from shared.utils import get_node_id_keypair
-from worker.download.download_utils import build_model_path
+from worker.download.impl_shard_downloader import exo_shard_downloader
+from worker.download.shard_downloader import RepoDownloadProgress, ShardDownloader
from worker.runner.runner_supervisor import RunnerSupervisor
from worker.utils.profile import start_polling_node_metrics
@@ -70,15 +72,15 @@ class AssignedRunner(BaseModel):
model_config = ConfigDict(arbitrary_types_allowed=True)
- @property
- def is_downloaded(self) -> bool:
- # TODO: Do this properly with huggingface validating each of the files.
- return os.path.exists(build_model_path(self.shard_metadata.model_meta.model_id))
-
+ is_downloaded: bool = False
+
+ def set_is_downloaded(self, is_downloaded: bool) -> None:
+ self.is_downloaded = is_downloaded
+
def status_update_event(self) -> RunnerStatusUpdated:
return RunnerStatusUpdated(
runner_id=self.runner_id,
- runner_status=self.status,
+ runner_status=deepcopy(self.status),
)
class Worker:
@@ -86,11 +88,13 @@ class Worker:
self,
node_id: NodeId,
logger: logging.Logger,
+ shard_downloader: ShardDownloader,
worker_events: AsyncSQLiteEventStorage | None,
global_events: AsyncSQLiteEventStorage | None,
):
self.node_id: NodeId = node_id
self.state: State = State()
+ self.shard_downloader: ShardDownloader = shard_downloader
self.worker_events: AsyncSQLiteEventStorage | None = worker_events # worker_events is None in some tests.
self.global_events: AsyncSQLiteEventStorage | None = global_events
self.logger: logging.Logger = logger
@@ -183,12 +187,26 @@ class Worker:
The model needs assigning and then downloading.
This op moves the runner from Assigned -> Downloading -> Ready state.
'''
+
+ initial_progress = await self.shard_downloader.get_shard_download_status_for_shard(op.shard_metadata)
+ if initial_progress.status == "complete":
+ self.assigned_runners[op.runner_id].set_is_downloaded(True)
+ self.assigned_runners[op.runner_id].status = DownloadingRunnerStatus(
+ download_progress=DownloadCompleted(
+ node_id=self.node_id,
+ )
+ )
+ yield self.assigned_runners[op.runner_id].status_update_event()
+ self.assigned_runners[op.runner_id].status = ReadyRunnerStatus()
+ yield self.assigned_runners[op.runner_id].status_update_event()
+ return
+
initial_status = DownloadingRunnerStatus(
download_progress=DownloadOngoing(
node_id=self.node_id,
download_progress=DownloadProgressData(
- total_bytes=1, # tmp
- downloaded_bytes=0
+ total_bytes=initial_progress.total_bytes,
+ downloaded_bytes=initial_progress.downloaded_bytes
)
)
)
@@ -206,25 +224,53 @@ class Worker:
# Download it!
# TODO: we probably want download progress as part of a callback that gets passed to the downloader.
+ download_progress_queue: asyncio.Queue[RepoDownloadProgress] = asyncio.Queue()
+ def download_progress_callback(shard: ShardMetadata, progress: RepoDownloadProgress) -> None:
+ download_progress_queue.put_nowait(progress)
- try:
- assert assigned_runner.is_downloaded
- assigned_runner.status = DownloadingRunnerStatus(
- download_progress=DownloadCompleted(
- node_id=self.node_id,
+
+ self.shard_downloader.on_progress(download_progress_callback)
+
+ asyncio.create_task(self.shard_downloader.ensure_shard(op.shard_metadata))
+
+ timeout_secs = 10 * 60
+ start_time = process_time()
+ last_yield_progress = start_time
+ while process_time() - start_time < timeout_secs:
+ progress: RepoDownloadProgress = await download_progress_queue.get()
+ if progress.status == "complete":
+ assigned_runner.status = DownloadingRunnerStatus(
+ download_progress=DownloadCompleted(
+ node_id=self.node_id,
+ )
)
- )
- except Exception as e:
+ yield assigned_runner.status_update_event()
+ assigned_runner.set_is_downloaded(True)
+ assigned_runner.status = ReadyRunnerStatus()
+ yield assigned_runner.status_update_event()
+ break
+ elif progress.status == "in_progress":
+ if process_time() - last_yield_progress > 1:
+ assigned_runner.status = DownloadingRunnerStatus(
+ download_progress=DownloadOngoing(
+ node_id=self.node_id,
+ download_progress=DownloadProgressData(
+ total_bytes=progress.total_bytes,
+ downloaded_bytes=progress.downloaded_bytes,
+ )
+ )
+ )
+ yield assigned_runner.status_update_event()
+ last_yield_progress = process_time()
+ else:
assigned_runner.status = DownloadingRunnerStatus(
download_progress=DownloadFailed(
node_id=self.node_id,
- error_message=str(e)
+ error_message=f"Timeout downloading model: {op.shard_metadata.model_meta.model_id}"
)
)
- yield assigned_runner.status_update_event()
+ yield assigned_runner.status_update_event()
- assigned_runner.status = ReadyRunnerStatus()
- yield assigned_runner.status_update_event()
async def _execute_task_op(
self, op: ExecuteTaskOp
@@ -383,6 +429,7 @@ class Worker:
for _instance_id, instance in state.instances.items():
if self.node_id in instance.shard_assignments.node_to_runner and \
instance.shard_assignments.node_to_runner[self.node_id] in state.runners and \
+ instance.shard_assignments.node_to_runner[self.node_id] in self.assigned_runners and \
isinstance(self.assigned_runners[instance.shard_assignments.node_to_runner[self.node_id]].status, FailedRunnerStatus):
num_spundown_nodes = 0
@@ -484,6 +531,7 @@ class Worker:
async def event_publisher(self, event: Event) -> None:
assert self.worker_events is not None
await self.worker_events.append_events([event], self.node_id)
+ print(f"published event: {event}")
# Handle state updates
async def run(self):
@@ -500,6 +548,8 @@ class Worker:
# 3. based on the updated state, we plan & execute an operation.
op: RunnerOp | None = self.plan(self.state)
+ if op is not None:
+ self.logger.info(f"!!! plan result: {op}")
# run the op, synchronously blocking for now
if op is not None:
@@ -507,7 +557,8 @@ class Worker:
await self.event_publisher(event)
await asyncio.sleep(0.01)
- self.logger.info(f"state: {self.state}")
+ if len(events) > 0:
+ self.logger.info(f"state: {self.state}")
async def main():
@@ -522,6 +573,7 @@ async def main():
event_log_manager = EventLogManager(EventLogConfig(), logger)
await event_log_manager.initialize()
+ shard_downloader = exo_shard_downloader()
# TODO: add profiling etc to resource monitor
async def resource_monitor_callback(node_performance_profile: NodePerformanceProfile) -> None:
@@ -530,7 +582,7 @@ async def main():
)
asyncio.create_task(start_polling_node_metrics(callback=resource_monitor_callback))
- worker = Worker(node_id, logger, event_log_manager.worker_events, event_log_manager.global_events)
+ worker = Worker(node_id, logger, shard_downloader, event_log_manager.worker_events, event_log_manager.global_events)
await worker.run()
diff --git a/worker/runner/runner_supervisor.py b/worker/runner/runner_supervisor.py
index 54d380d2..43b515dc 100644
--- a/worker/runner/runner_supervisor.py
+++ b/worker/runner/runner_supervisor.py
@@ -197,7 +197,7 @@ class RunnerSupervisor:
text=text, token=token, finish_reason=finish_reason
):
yield TokenChunk(
- command_id=CommandId(task.task_id),
+ command_id=CommandId(task.command_id),
idx=token,
model=self.model_shard_meta.model_meta.model_id,
text=text,
diff --git a/worker/tests/conftest.py b/worker/tests/conftest.py
index ad76fdab..9ef65c3d 100644
--- a/worker/tests/conftest.py
+++ b/worker/tests/conftest.py
@@ -9,7 +9,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 NodeId
+from shared.types.common import CommandId, NodeId
from shared.types.models import ModelId, ModelMetadata
from shared.types.state import State
from shared.types.tasks import (
@@ -27,6 +27,7 @@ from shared.types.worker.ops import (
)
from shared.types.worker.runners import RunnerId, ShardAssignments
from shared.types.worker.shards import PipelineShardMetadata
+from worker.download.shard_downloader import NoopShardDownloader
from worker.main import Worker
@@ -103,6 +104,7 @@ def chat_completion_task(completion_create_params: ChatCompletionTaskParams):
def _chat_completion_task(instance_id: InstanceId) -> ChatCompletionTask:
return ChatCompletionTask(
task_id=TaskId(),
+ command_id=CommandId(),
instance_id=instance_id,
task_type=TaskType.CHAT_COMPLETION,
task_status=TaskStatus.PENDING,
@@ -153,9 +155,10 @@ def instance(pipeline_shard_meta: Callable[[int, int], PipelineShardMetadata], h
@pytest.fixture
async def worker(node_id: NodeId, logger: Logger):
event_log_manager = EventLogManager(EventLogConfig(), logger)
+ shard_downloader = NoopShardDownloader()
await event_log_manager.initialize()
- return Worker(node_id, logger, worker_events=event_log_manager.global_events, global_events=event_log_manager.global_events)
+ return Worker(node_id, logger, shard_downloader, worker_events=event_log_manager.global_events, global_events=event_log_manager.global_events)
@pytest.fixture
async def worker_with_assigned_runner(worker: Worker, instance: Callable[[InstanceId, NodeId, RunnerId], Instance]):
@@ -204,7 +207,8 @@ def worker_running(logger: Logger) -> Callable[[NodeId], Awaitable[tuple[Worker,
global_events = event_log_manager.global_events
await global_events.delete_all_events()
- worker = Worker(node_id, logger=logger, worker_events=global_events, global_events=global_events)
+ shard_downloader = NoopShardDownloader()
+ worker = Worker(node_id, logger=logger, shard_downloader=shard_downloader, worker_events=global_events, global_events=global_events)
asyncio.create_task(worker.run())
return worker, global_events
diff --git a/worker/tests/test_worker_integration.py b/worker/tests/test_worker_integration.py
index 3041080c..acd28735 100644
--- a/worker/tests/test_worker_integration.py
+++ b/worker/tests/test_worker_integration.py
@@ -32,6 +32,7 @@ from shared.types.worker.runners import (
# RunningRunnerStatus,
)
from shared.types.worker.shards import PipelineShardMetadata
+from worker.download.shard_downloader import NoopShardDownloader
from worker.main import AssignedRunner, Worker
from worker.tests.test_worker_integration_utils import read_streaming_response
@@ -269,14 +270,15 @@ async def test_2_runner_inference(
):
event_log_manager = EventLogManager(EventLogConfig(), logger)
await event_log_manager.initialize()
+ shard_downloader = NoopShardDownloader()
global_events = event_log_manager.global_events
await global_events.delete_all_events()
- worker1 = Worker(NODE_A, logger=logger, worker_events=global_events, global_events=global_events)
+ worker1 = Worker(NODE_A, logger=logger, shard_downloader=shard_downloader, worker_events=global_events, global_events=global_events)
asyncio.create_task(worker1.run())
- worker2 = Worker(NODE_B, logger=logger, worker_events=global_events, global_events=global_events)
+ worker2 = Worker(NODE_B, logger=logger, shard_downloader=shard_downloader, worker_events=global_events, global_events=global_events)
asyncio.create_task(worker2.run())
## Instance
@@ -348,14 +350,15 @@ async def test_runner_respawn(
):
event_log_manager = EventLogManager(EventLogConfig(), logger)
await event_log_manager.initialize()
+ shard_downloader = NoopShardDownloader()
global_events = event_log_manager.global_events
await global_events.delete_all_events()
- worker1 = Worker(NODE_A, logger=logger, worker_events=global_events, global_events=global_events)
+ worker1 = Worker(NODE_A, logger=logger, shard_downloader=shard_downloader, worker_events=global_events, global_events=global_events)
asyncio.create_task(worker1.run())
- worker2 = Worker(NODE_B, logger=logger, worker_events=global_events, global_events=global_events)
+ worker2 = Worker(NODE_B, logger=logger, shard_downloader=shard_downloader, worker_events=global_events, global_events=global_events)
asyncio.create_task(worker2.run())
## Instance
diff --git a/worker/tests/test_worker_plan.py b/worker/tests/test_worker_plan.py
index 120e3895..040d47ee 100644
--- a/worker/tests/test_worker_plan.py
+++ b/worker/tests/test_worker_plan.py
@@ -36,9 +36,11 @@ from shared.types.worker.runners import (
)
from shared.types.worker.shards import PipelineShardMetadata
from worker.download.download_utils import build_model_path
-from worker.main import Worker
+from worker.download.shard_downloader import NoopShardDownloader
+from worker.main import AssignedRunner, Worker
from .test_worker_plan_utils import (
+ COMMAND_1_ID,
INSTANCE_1_ID,
MODEL_A_ID,
NODE_A,
@@ -47,7 +49,6 @@ from .test_worker_plan_utils import (
RUNNER_2_ID,
TASK_1_ID,
InProcessRunner,
- OverrideAssignedRunner,
PlanTestCase,
make_downloading_status,
make_model_meta,
@@ -339,7 +340,7 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
)
},
runners={RUNNER_1_ID: ReadyRunnerStatus(), RUNNER_2_ID: DownloadingRunnerStatus(download_progress=DownloadPending(node_id=NODE_A))},
- tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
+ tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, command_id=COMMAND_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
),
expected_op=None
),
@@ -382,7 +383,7 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
)
},
runners={RUNNER_1_ID: ReadyRunnerStatus(), RUNNER_2_ID: ReadyRunnerStatus()},
- tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
+ tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, command_id=COMMAND_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
),
expected_op=RunnerUpOp(runner_id=RUNNER_1_ID)
),
@@ -484,6 +485,7 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
tasks={
TASK_1_ID: ChatCompletionTask(
task_id=TASK_1_ID,
+ command_id=COMMAND_1_ID,
task_type=TaskType.CHAT_COMPLETION,
task_status=TaskStatus.PENDING,
task_params=ChatCompletionTaskParams(
@@ -501,6 +503,7 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
),
expected_op=ExecuteTaskOp(runner_id=RUNNER_1_ID, task=ChatCompletionTask(
task_id=TASK_1_ID,
+ command_id=COMMAND_1_ID,
instance_id=INSTANCE_1_ID,
task_type=TaskType.CHAT_COMPLETION,
task_status=TaskStatus.PENDING,
@@ -550,7 +553,7 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
)
},
runners={RUNNER_1_ID: LoadedRunnerStatus(), RUNNER_2_ID: LoadedRunnerStatus()},
- tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
+ tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, command_id=COMMAND_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
),
expected_op=None
),
@@ -593,12 +596,13 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
)
},
runners={RUNNER_1_ID: LoadedRunnerStatus(), RUNNER_2_ID: LoadedRunnerStatus()},
- tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
+ tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, command_id=COMMAND_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
),
expected_op=ExecuteTaskOp(
runner_id=RUNNER_1_ID,
task=ChatCompletionTask(
task_id=TASK_1_ID,
+ command_id=COMMAND_1_ID,
instance_id=INSTANCE_1_ID,
task_type=TaskType.CHAT_COMPLETION,
task_params=ChatCompletionTaskParams(
@@ -648,12 +652,13 @@ def _get_test_cases(tmp_path: Path) -> list[PlanTestCase]:
)
},
runners={RUNNER_1_ID: LoadedRunnerStatus(), RUNNER_2_ID: RunningRunnerStatus()},
- tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
+ tasks={TASK_1_ID: ChatCompletionTask(task_id=TASK_1_ID, command_id=COMMAND_1_ID, task_type=TaskType.CHAT_COMPLETION, task_status=TaskStatus.PENDING, task_params=ChatCompletionTaskParams(model=str(MODEL_A_ID), messages=[ChatCompletionMessage(role="user", content="Hello, world!")]), instance_id=INSTANCE_1_ID)},
),
expected_op=ExecuteTaskOp(
runner_id=RUNNER_1_ID,
task=ChatCompletionTask(
task_id=TASK_1_ID,
+ command_id=COMMAND_1_ID,
instance_id=INSTANCE_1_ID,
task_type=TaskType.CHAT_COMPLETION,
task_params=ChatCompletionTaskParams(
@@ -851,7 +856,8 @@ def test_worker_plan(case: PlanTestCase, tmp_path: Path, monkeypatch: pytest.Mon
node_id = NODE_A
logger = logging.getLogger("test_worker_plan")
- worker = Worker(node_id=node_id, worker_events=None, global_events=None, logger=logger)
+ shard_downloader = NoopShardDownloader()
+ worker = Worker(node_id=node_id, shard_downloader=shard_downloader, worker_events=None, global_events=None, logger=logger)
path_downloaded_map: dict[str, bool] = {}
@@ -891,25 +897,17 @@ def test_worker_plan(case: PlanTestCase, tmp_path: Path, monkeypatch: pytest.Mon
raise Exception('test_worker_plan not currently designed to have more than 1 instance.')
- assigned_runner = OverrideAssignedRunner(
+ assigned_runner = AssignedRunner(
runner_id=runner_config.runner_id,
instance_id=runner_config.instance_id,
shard_metadata=shard_metadata,
hosts=[],
status=runner_config.status,
runner=None,
- downloaded=runner_config.downloaded
+ is_downloaded=runner_config.downloaded
)
worker.assigned_runners[runner_config.runner_id] = assigned_runner
path_downloaded_map[str(build_model_path(shard_metadata.model_meta.model_id))] = runner_config.downloaded
- # Stub filesystem existence check ------------------------------------------------------
- from worker import main as worker_main # local import for module-scoped os
-
- def _fake_exists(path: str | Path) -> bool: # noqa: ANN001 – match os.path.exists signature
- return path_downloaded_map.get(str(path), False)
-
- monkeypatch.setattr(worker_main.os.path, "exists", _fake_exists)
-
op = worker.plan(case.state)
assert op == case.expected_op
diff --git a/worker/tests/test_worker_plan_utils.py b/worker/tests/test_worker_plan_utils.py
index b0c81fad..84d92ab0 100644
--- a/worker/tests/test_worker_plan_utils.py
+++ b/worker/tests/test_worker_plan_utils.py
@@ -2,10 +2,10 @@ from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
-from typing import Final, List, Optional, override
+from typing import Final, List, Optional
from shared.models.model_cards import MODEL_CARDS, ModelCard
-from shared.types.common import NodeId
+from shared.types.common import CommandId, NodeId
from shared.types.models import ModelId, ModelMetadata
from shared.types.state import State
from shared.types.tasks import TaskId
@@ -20,7 +20,6 @@ from shared.types.worker.runners import (
ShardAssignments,
)
from shared.types.worker.shards import PipelineShardMetadata
-from worker.main import AssignedRunner
NODE_A: Final[NodeId] = NodeId("aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa")
NODE_B: Final[NodeId] = NodeId("bbbbbbbb-bbbb-4bbb-8bbb-bbbbbbbbbbbb")
@@ -33,11 +32,11 @@ INSTANCE_2_ID: Final[InstanceId] = InstanceId()
MODEL_A_ID: Final[ModelId] = 'mlx-community/Llama-3.2-1B-Instruct-4bit'
MODEL_B_ID: Final[ModelId] = 'mlx-community/Llama-3.2-1B-Instruct-4bit'
TASK_1_ID: Final[TaskId] = TaskId()
+COMMAND_1_ID: Final[CommandId] = CommandId()
@dataclass(slots=True, frozen=True)
class InProcessRunner:
"""Minimal description of a runner's in-process state."""
- # TODO: Rename to InProcessRunnerConfig and create a constructor for OverrideAssignedRunner.
runner_id: RunnerId
instance_id: InstanceId
@@ -46,15 +45,6 @@ class InProcessRunner:
downloaded: bool
device_rank: int = 0
-# Helper class to override the is_downloaded property to whatever is specified by InProcessRunner
-class OverrideAssignedRunner(AssignedRunner):
- downloaded: bool
-
- @property
- @override
- def is_downloaded(self) -> bool:
- return self.downloaded
-
@dataclass(slots=True, frozen=True)
class PlanTestCase:
← b687dec6 Discovery integration master
·
back to Exo
·
fix placement tests b285a9f0 →