← back to Exo
Glue TWO
a97fb27c64e7cc804f2069904e5a4a861f86c4a7 · 2025-07-25 14:32:34 +0100 · Alex Cheema
Files touched
M master/main.pyM shared/constants.pyM worker/main.pyM worker/tests/conftest.pyM worker/tests/test_worker_plan.py
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 →