← back to Exo
fix api get_state
2031d9481d16798216a707fcbf3dabb8d2c23531 · 2025-07-30 07:15:15 -0700 · Alex Cheema
Files touched
M master/api.pyM master/tests/test_master.pyM worker/main.pyM worker/tests/test_supervisor_errors.py
Diff
commit 2031d9481d16798216a707fcbf3dabb8d2c23531
Author: Alex Cheema <41707476+AlexCheema@users.noreply.github.com>
Date: Wed Jul 30 07:15:15 2025 -0700
fix api get_state
---
master/api.py | 13 +++++--------
master/tests/test_master.py | 23 +++++++++++++++++++----
worker/main.py | 4 ++--
worker/tests/test_supervisor_errors.py | 4 ++--
4 files changed, 28 insertions(+), 16 deletions(-)
diff --git a/master/api.py b/master/api.py
index ba74077f..d6f1a091 100644
--- a/master/api.py
+++ b/master/api.py
@@ -1,7 +1,7 @@
import asyncio
-from pathlib import Path
import time
from collections.abc import AsyncGenerator
+from pathlib import Path
from typing import Callable, List, Sequence, final
import uvicorn
@@ -72,14 +72,14 @@ async def resolve_model_meta(model_id: str) -> ModelMetadata:
@final
class API:
def __init__(self, command_buffer: List[Command], global_events: AsyncSQLiteEventStorage, get_state: Callable[[], State]) -> None:
+ self.get_state = get_state
+ self.command_buffer = command_buffer
+ self.global_events = global_events
+
self._app = FastAPI()
self._setup_cors()
self._setup_routes()
- self.command_buffer = command_buffer
- self.global_events = global_events
- self.get_state = get_state
-
self._app.mount("/", StaticFiles(directory=_DASHBOARD_DIR, html=True), name="dashboard")
def _setup_cors(self) -> None:
@@ -208,9 +208,6 @@ class API:
description=card.description,
tags=card.tags) for card in MODEL_CARDS.values()])
- async def get_state(self) -> State:
- return self.get_state()
-
def start_fastapi_server(
command_buffer: List[Command],
diff --git a/master/tests/test_master.py b/master/tests/test_master.py
index 14125987..a6649495 100644
--- a/master/tests/test_master.py
+++ b/master/tests/test_master.py
@@ -2,7 +2,7 @@ import asyncio
import tempfile
from logging import Logger
from pathlib import Path
-from typing import List
+from typing import List, Sequence
import pytest
from exo_pyo3_bindings import Keypair
@@ -13,7 +13,7 @@ from shared.db.sqlite.connector import AsyncSQLiteEventStorage
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 import Event, EventFromEventLog, Heartbeat, TaskCreated
from shared.types.events._events import (
InstanceCreated,
NodePerformanceMeasured,
@@ -53,6 +53,21 @@ async def test_master():
global_events: AsyncSQLiteEventStorage = event_log_manager.global_events
await global_events.delete_all_events()
+ async def _get_events() -> Sequence[EventFromEventLog[Event]]:
+ orig_events = await global_events.get_events_since(0)
+ override_idx_in_log = 1
+ events: List[EventFromEventLog[Event]] = []
+ for e in orig_events:
+ if isinstance(e.event, Heartbeat):
+ continue
+ events.append(EventFromEventLog(
+ event=e.event,
+ origin=e.origin,
+ idx_in_log=override_idx_in_log
+ ))
+ override_idx_in_log += 1
+ return events
+
command_buffer: List[Command] = []
forwarder_binary_path = _create_forwarder_dummy_binary()
@@ -104,10 +119,10 @@ async def test_master():
)
)
)
- while len(await global_events.get_events_since(0)) < 4:
+ while len(await _get_events()) < 4:
await asyncio.sleep(0.001)
- events = await global_events.get_events_since(0)
+ events = await _get_events()
print(events)
assert len(events) == 4
assert events[0].idx_in_log == 1
diff --git a/worker/main.py b/worker/main.py
index bf537302..0fd25765 100644
--- a/worker/main.py
+++ b/worker/main.py
@@ -151,12 +151,12 @@ class Worker:
# TODO: This should be dynamic, based on the size of the model.
if not initialize_timeout:
- GBPS = 10
+ gigabytes_per_second = 10
shard = assigned_runner.shard_metadata
weights_size_kb = (shard.end_layer - shard.start_layer) / shard.n_layers * shard.model_meta.storage_size_kilobytes
- initialize_timeout = weights_size_kb / (1024**2 * GBPS) + 2.0 # Add a constant 2.0 to ensure connection can be made as well
+ initialize_timeout = weights_size_kb / (1024**2 * gigabytes_per_second) + 2.0 # Add a constant 2.0 to ensure connection can be made as well
try:
assigned_runner.runner = await asyncio.wait_for(
diff --git a/worker/tests/test_supervisor_errors.py b/worker/tests/test_supervisor_errors.py
index 8b13ef62..87390898 100644
--- a/worker/tests/test_supervisor_errors.py
+++ b/worker/tests/test_supervisor_errors.py
@@ -15,8 +15,8 @@ from shared.types.events import (
InstanceDeleted,
RunnerStatusUpdated,
TaskCreated,
- TaskStateUpdated,
TaskFailed,
+ TaskStateUpdated,
)
from shared.types.events.chunks import GenerationChunk, TokenChunk
from shared.types.models import ModelId
@@ -57,7 +57,7 @@ async def test_stream_response_failed_always(
instance: Callable[[InstanceId, NodeId, RunnerId], Instance],
chat_completion_task: Callable[[InstanceId, TaskId], Task]
):
- worker, global_events = await worker_running(NODE_A)
+ _, global_events = await worker_running(NODE_A)
instance_value: Instance = instance(INSTANCE_1_ID, NODE_A, RUNNER_1_ID)
instance_value.instance_type = InstanceStatus.ACTIVE
← b350eded Test Supervisor Errors.
·
back to Exo
·
fix libp2p + other prs that were wrongly overwritten before 0e32599e →