[object Object]

← back to Exo

fix api get_state

2031d9481d16798216a707fcbf3dabb8d2c23531 · 2025-07-30 07:15:15 -0700 · Alex Cheema

Files touched

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 →