[object Object]

← back to Exo

Glue TWO

a97fb27c64e7cc804f2069904e5a4a861f86c4a7 · 2025-07-25 14:32:34 +0100 · Alex Cheema

Files touched

Diff

commit a97fb27c64e7cc804f2069904e5a4a861f86c4a7
Author: Alex Cheema <41707476+AlexCheema@users.noreply.github.com>
Date:   Fri Jul 25 14:32:34 2025 +0100

    Glue TWO
---
 master/main.py                   | 50 ++++++++++++++--------------------------
 shared/constants.py              |  4 +++-
 worker/main.py                   | 29 ++++++++++++++---------
 worker/tests/conftest.py         |  4 ++--
 worker/tests/test_worker_plan.py |  2 +-
 5 files changed, 41 insertions(+), 48 deletions(-)

diff --git a/master/main.py b/master/main.py
index acc1b122..3c1e8a57 100644
--- a/master/main.py
+++ b/master/main.py
@@ -1,7 +1,8 @@
 import asyncio
+import logging
 import os
 import threading
-from logging import Logger
+import traceback
 from pathlib import Path
 from typing import List
 
@@ -17,7 +18,6 @@ from shared.node_id import get_node_id_keypair
 from shared.types.common import NodeId
 from shared.types.events import (
     Event,
-    NodePerformanceMeasured,
     TaskCreated,
 )
 from shared.types.events.commands import (
@@ -26,18 +26,13 @@ from shared.types.events.commands import (
     CreateInstanceCommand,
     DeleteInstanceCommand,
 )
-from shared.types.profiling import (
-    MemoryPerformanceProfile,
-    NodePerformanceProfile,
-    SystemPerformanceProfile,
-)
 from shared.types.state import State
 from shared.types.tasks import ChatCompletionTask, TaskId, TaskStatus, TaskType
 from shared.types.worker.instances import Instance
 
 
 class Master:
-    def __init__(self, node_id: NodeId, command_buffer: list[Command], global_events: AsyncSQLiteEventStorage, forwarder_binary_path: Path, logger: Logger):
+    def __init__(self, node_id: NodeId, command_buffer: list[Command], global_events: AsyncSQLiteEventStorage, forwarder_binary_path: Path, logger: logging.Logger):
         self.node_id = node_id
         self.command_buffer = command_buffer
         self.global_events = global_events
@@ -53,13 +48,9 @@ class Master:
         return State()
 
     async def _run_event_loop_body(self) -> None:
-        if self.forwarder_supervisor.current_role == ForwarderRole.REPLICA:
-            await asyncio.sleep(0.1)
-            return
-
         next_events: list[Event] = []
         # 1. process commands
-        if len(self.command_buffer) > 0:
+        if self.forwarder_supervisor.current_role == ForwarderRole.MASTER and len(self.command_buffer) > 0:
             # for now we do one command at a time
             next_command = self.command_buffer.pop(0)
             self.logger.info(f"got command: {next_command}")
@@ -106,7 +97,7 @@ class Master:
         for event_from_log in events:
             self.state = apply(self.state, event_from_log)
 
-        self.logger.info(f"state: {self.state.model_dump_json()}")
+        self.logger.info(f"state: {self.state}")
 
     async def run(self):
         self.state = await self._get_state_snapshot()
@@ -123,12 +114,19 @@ class Master:
                 await self._run_event_loop_body()
             except Exception as e:
                 self.logger.error(f"Error in _run_event_loop_body: {e}")
+                traceback.print_exc()
                 await asyncio.sleep(0.1)
 
 
 
 async def main():
-    logger = Logger(name='master_logger')
+    logger = logging.getLogger('master_logger')
+    logger.setLevel(logging.DEBUG)
+    if not logger.handlers:
+        handler = logging.StreamHandler()
+        handler.setFormatter(logging.Formatter('%(asctime)s - %(levelname)s - %(message)s'))
+        logger.addHandler(handler)
+
     node_id_keypair = get_node_id_keypair()
     node_id = NodeId(node_id_keypair.to_peer_id().to_base58())
 
@@ -136,32 +134,18 @@ async def main():
     await event_log_manager.initialize()
     global_events: AsyncSQLiteEventStorage = event_log_manager.global_events
 
-    # TODO: this should be the resource monitor that does this
-    await global_events.append_events([NodePerformanceMeasured(
-        node_id=node_id,
-        node_profile=NodePerformanceProfile(
-            model_id="testmodelabc",
-            chip_id="testchipabc",
-            memory=MemoryPerformanceProfile(
-                ram_total=1000,
-                ram_available=1000,
-                swap_total=1000,
-                swap_available=1000
-            ),
-            system=SystemPerformanceProfile(
-                flops_fp16=1000
-            )
-        )
-    )], origin=node_id)
-
     command_buffer: List[Command] = []
 
+    logger.info(f"Starting Master with node_id: {node_id}")
+
     api_thread = threading.Thread(
         target=start_fastapi_server,
         args=(
             command_buffer,
             global_events,
             lambda: master.state,
+            "0.0.0.0",
+            int(os.environ.get("API_PORT", 8000))
         ),
         daemon=True
     )
diff --git a/shared/constants.py b/shared/constants.py
index 61119538..6f30ab88 100644
--- a/shared/constants.py
+++ b/shared/constants.py
@@ -1,7 +1,9 @@
 import inspect
+import os
 from pathlib import Path
 
-EXO_HOME = Path.home() / ".exo"
+EXO_HOME_RELATIVE_PATH = os.environ.get("EXO_HOME", ".exo")
+EXO_HOME = Path.home() / EXO_HOME_RELATIVE_PATH
 EXO_GLOBAL_EVENT_DB = EXO_HOME / "global_events.db"
 EXO_WORKER_EVENT_DB = EXO_HOME / "worker_events.db"
 EXO_MASTER_STATE = EXO_HOME / "master_state.json"
diff --git a/worker/main.py b/worker/main.py
index 3c4a5c45..0196116c 100644
--- a/worker/main.py
+++ b/worker/main.py
@@ -1,8 +1,8 @@
 import asyncio
+import logging
 import os
 from asyncio import Queue
 from functools import partial
-from logging import Logger
 from typing import AsyncGenerator, Optional
 
 from pydantic import BaseModel, ConfigDict
@@ -10,6 +10,7 @@ from pydantic import BaseModel, ConfigDict
 from shared.apply import apply
 from shared.db.sqlite import AsyncSQLiteEventStorage
 from shared.db.sqlite.event_log_manager import EventLogConfig, EventLogManager
+from shared.node_id import get_node_id_keypair
 from shared.types.common import NodeId
 from shared.types.events import (
     ChunkGenerated,
@@ -57,9 +58,6 @@ from worker.runner.runner_supervisor import RunnerSupervisor
 from worker.utils.profile import start_polling_node_metrics
 
 
-def get_node_id() -> NodeId:
-    return NodeId() # TODO
-
 class AssignedRunner(BaseModel):
     runner_id: RunnerId
     instance_id: InstanceId
@@ -86,13 +84,15 @@ class Worker:
     def __init__(
         self,
         node_id: NodeId,
-        logger: Logger,
+        logger: logging.Logger,
         worker_events: AsyncSQLiteEventStorage | None,
+        global_events: AsyncSQLiteEventStorage | None,
     ):
         self.node_id: NodeId = node_id
         self.state: State = State()
         self.worker_events: AsyncSQLiteEventStorage | None = worker_events # worker_events is None in some tests.
-        self.logger: Logger = logger
+        self.global_events: AsyncSQLiteEventStorage | None = global_events
+        self.logger: logging.Logger = logger
 
         self.assigned_runners: dict[RunnerId, AssignedRunner] = {}
         self._task: asyncio.Task[None] | None = None
@@ -462,11 +462,11 @@ class Worker:
 
     # Handle state updates
     async def run(self):
-        assert self.worker_events is not None
+        assert self.global_events is not None
 
         while True:
             # 1. get latest events
-            events = await self.worker_events.get_events_since(self.state.last_event_applied_idx)
+            events = await self.global_events.get_events_since(self.state.last_event_applied_idx)
             if len(events) == 0:
                 await asyncio.sleep(0.01)
                 continue
@@ -484,11 +484,18 @@ class Worker:
                     await self.event_publisher(event)
 
             await asyncio.sleep(0.01)
+            self.logger.info(f"state: {self.state}")
 
 
 async def main():
-    node_id: NodeId = get_node_id()
-    logger: Logger = Logger('worker_log')
+    node_id_keypair = get_node_id_keypair()
+    node_id = NodeId(node_id_keypair.to_peer_id().to_base58())
+    logger: logging.Logger = logging.getLogger('worker_logger')
+    logger.setLevel(logging.DEBUG)
+    if not logger.handlers:
+        handler = logging.StreamHandler()
+        handler.setFormatter(logging.Formatter('%(asctime)s - %(levelname)s - %(message)s'))
+        logger.addHandler(handler)
 
     event_log_manager = EventLogManager(EventLogConfig(), logger)
     await event_log_manager.initialize()
@@ -500,7 +507,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)
+    worker = Worker(node_id, logger, event_log_manager.worker_events, event_log_manager.global_events)
 
     await worker.run()
 
diff --git a/worker/tests/conftest.py b/worker/tests/conftest.py
index 70f230b2..38ed90d8 100644
--- a/worker/tests/conftest.py
+++ b/worker/tests/conftest.py
@@ -153,7 +153,7 @@ async def worker(node_id: NodeId, logger: Logger):
     event_log_manager = EventLogManager(EventLogConfig(), logger)
     await event_log_manager.initialize()
 
-    return Worker(node_id, logger, worker_events=event_log_manager.global_events)
+    return Worker(node_id, logger, 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[[NodeId, RunnerId], Instance]):
@@ -202,7 +202,7 @@ 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)
+        worker = Worker(node_id, logger=logger, worker_events=global_events, global_events=global_events)
         asyncio.create_task(worker.run())
 
         return worker, global_events
diff --git a/worker/tests/test_worker_plan.py b/worker/tests/test_worker_plan.py
index 3da7c8c8..8f00b84b 100644
--- a/worker/tests/test_worker_plan.py
+++ b/worker/tests/test_worker_plan.py
@@ -835,7 +835,7 @@ 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, logger=logger)
+    worker = Worker(node_id=node_id, worker_events=None, global_events=None, logger=logger)
 
     path_downloaded_map: dict[str, bool] = {}
 

← 9be08ec7 add resource monitor  ·  back to Exo  ·  Serialize topology 261e5752 →