← back to Exo
Sync custom model cards across nodes (#1768)
49951e1b1a7c6351d2013aa5c7e6d3aae2b1e41b · 2026-03-24 12:51:36 +0000 · ciaranbor
## Motivation
Custom model cards were only saved locally on the node that handled the
API request.
## Changes
- Added AddCustomModelCard and DeleteCustomModelCard commands
- Added CustomModelCardAdded and CustomModelCardDeleted events
- Added custom_model_cards field to cluster State
- Master handles new commands by emitting corresponding events
- Workers persist model cards to disk and update the in-memory cache on
event receipt
- Separated fetch_from_hf (pure fetch) from disk persistence (now
handled by event-sourcing layer)
- Exposed add_to_card_cache() helper for the worker to update the cache
## Why It Works
- Follows the existing event-sourcing pattern: API → Command → Master →
Event → all Workers
- Every node applies the same events, so custom model cards are
consistent across the cluster
## Test Plan
### Manual Testing
Add/delete a custom model via the API on one node, verify it
appears/disappears on all nodes
### Automated Testing
Existing tests cover apply() logic; new event types follow the same
discriminated-union pattern
Files touched
M src/exo/api/main.pyM src/exo/master/main.pyM src/exo/shared/apply.pyM src/exo/shared/models/model_cards.pyM src/exo/shared/types/commands.pyM src/exo/shared/types/events.pyM src/exo/worker/main.py
Diff
commit 49951e1b1a7c6351d2013aa5c7e6d3aae2b1e41b
Author: ciaranbor <81697641+ciaranbor@users.noreply.github.com>
Date: Tue Mar 24 12:51:36 2026 +0000
Sync custom model cards across nodes (#1768)
## Motivation
Custom model cards were only saved locally on the node that handled the
API request.
## Changes
- Added AddCustomModelCard and DeleteCustomModelCard commands
- Added CustomModelCardAdded and CustomModelCardDeleted events
- Added custom_model_cards field to cluster State
- Master handles new commands by emitting corresponding events
- Workers persist model cards to disk and update the in-memory cache on
event receipt
- Separated fetch_from_hf (pure fetch) from disk persistence (now
handled by event-sourcing layer)
- Exposed add_to_card_cache() helper for the worker to update the cache
## Why It Works
- Follows the existing event-sourcing pattern: API → Command → Master →
Event → all Workers
- Every node applies the same events, so custom model cards are
consistent across the cluster
## Test Plan
### Manual Testing
Add/delete a custom model via the API on one node, verify it
appears/disappears on all nodes
### Automated Testing
Existing tests cover apply() logic; new event types follow the same
discriminated-union pattern
---
src/exo/api/main.py | 30 ++++++++++++----
src/exo/master/main.py | 12 +++++++
src/exo/shared/apply.py | 4 +++
src/exo/shared/models/model_cards.py | 69 +++++++++++++++++++++---------------
src/exo/shared/types/commands.py | 10 ++++++
src/exo/shared/types/events.py | 13 ++++++-
src/exo/worker/main.py | 11 +++++-
7 files changed, 112 insertions(+), 37 deletions(-)
diff --git a/src/exo/api/main.py b/src/exo/api/main.py
index 08d95880..32900c78 100644
--- a/src/exo/api/main.py
+++ b/src/exo/api/main.py
@@ -129,9 +129,8 @@ from exo.shared.logging import InterceptLogger
from exo.shared.models.model_cards import (
ModelCard,
ModelId,
- delete_custom_card,
+ get_card,
get_model_cards,
- is_custom_card,
)
from exo.shared.tracing import TraceEvent, compute_stats, export_trace, load_trace_file
from exo.shared.types.chunks import (
@@ -143,8 +142,10 @@ from exo.shared.types.chunks import (
ToolCallChunk,
)
from exo.shared.types.commands import (
+ AddCustomModelCard,
Command,
CreateInstance,
+ DeleteCustomModelCard,
DeleteDownload,
DeleteInstance,
DownloadCommand,
@@ -1558,7 +1559,7 @@ class API:
storage_size_megabytes=card.storage_size.in_mb,
supports_tensor=card.supports_tensor,
tasks=[task.value for task in card.tasks],
- is_custom=is_custom_card(card.model_id),
+ is_custom=card.is_custom,
family=card.family,
quantization=card.quantization,
base_model=card.base_model,
@@ -1569,7 +1570,7 @@ class API:
)
async def add_custom_model(self, payload: AddCustomModelParams) -> ModelListModel:
- """Fetch a model from HuggingFace and save as a custom model card."""
+ """Fetch a model from HuggingFace and save as a custom model card, then sync across the cluster."""
try:
card = await ModelCard.fetch_from_hf(payload.model_id)
except Exception as exc:
@@ -1577,6 +1578,13 @@ class API:
status_code=400, detail=f"Failed to fetch model: {exc}"
) from exc
+ await self.command_sender.send(
+ ForwarderCommand(
+ origin=self._system_id,
+ command=AddCustomModelCard(model_card=card),
+ )
+ )
+
return ModelListModel(
id=card.model_id,
hugging_face_id=card.model_id,
@@ -1590,10 +1598,18 @@ class API:
)
async def delete_custom_model(self, model_id: ModelId) -> JSONResponse:
- """Delete a user-added custom model card."""
- deleted = await delete_custom_card(model_id)
- if not deleted:
+ """Delete a user-added custom model card and sync deletion across the cluster."""
+ card = get_card(model_id)
+ if card is None or not card.is_custom:
raise HTTPException(status_code=404, detail="Custom model card not found")
+
+ await self.command_sender.send(
+ ForwarderCommand(
+ origin=self._system_id,
+ command=DeleteCustomModelCard(model_id=model_id),
+ )
+ )
+
return JSONResponse(
{"message": "Model card deleted", "model_id": str(model_id)}
)
diff --git a/src/exo/master/main.py b/src/exo/master/main.py
index 94701c6c..1511897e 100644
--- a/src/exo/master/main.py
+++ b/src/exo/master/main.py
@@ -13,7 +13,9 @@ from exo.master.placement import (
from exo.shared.apply import apply
from exo.shared.constants import EXO_EVENT_LOG_DIR, EXO_TRACING_ENABLED
from exo.shared.types.commands import (
+ AddCustomModelCard,
CreateInstance,
+ DeleteCustomModelCard,
DeleteInstance,
ForwarderCommand,
ForwarderDownloadCommand,
@@ -29,6 +31,8 @@ from exo.shared.types.commands import (
)
from exo.shared.types.common import CommandId, NodeId, SessionId, SystemId
from exo.shared.types.events import (
+ CustomModelCardAdded,
+ CustomModelCardDeleted,
Event,
GlobalForwarderEvent,
IndexedEvent,
@@ -345,6 +349,14 @@ class Master:
f"Finished command {command.finished_command_id} finished"
)
+ case AddCustomModelCard():
+ generated_events.append(
+ CustomModelCardAdded(model_card=command.model_card)
+ )
+ case DeleteCustomModelCard():
+ generated_events.append(
+ CustomModelCardDeleted(model_id=command.model_id)
+ )
case RequestEventLog():
# We should just be able to send everything, since other buffers will ignore old messages
# rate limit to 1000 at a time
diff --git a/src/exo/shared/apply.py b/src/exo/shared/apply.py
index 38132a32..fb9f6d9a 100644
--- a/src/exo/shared/apply.py
+++ b/src/exo/shared/apply.py
@@ -7,6 +7,8 @@ from loguru import logger
from exo.shared.types.common import NodeId
from exo.shared.types.events import (
ChunkGenerated,
+ CustomModelCardAdded,
+ CustomModelCardDeleted,
Event,
IndexedEvent,
InputChunkReceived,
@@ -65,6 +67,8 @@ def event_apply(event: Event, state: State) -> State:
| InputChunkReceived()
| TracesCollected()
| TracesMerged()
+ | CustomModelCardAdded()
+ | CustomModelCardDeleted()
): # Pass-through events that don't modify state
return state
case InstanceCreated():
diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py
index 048c469f..fef1efaa 100644
--- a/src/exo/shared/models/model_cards.py
+++ b/src/exo/shared/models/model_cards.py
@@ -30,30 +30,42 @@ from exo.utils.pydantic_ext import CamelCaseModel
# kinda ugly...
# TODO: load search path from config.toml
_custom_cards_dir = Path(str(EXO_CUSTOM_MODEL_CARDS_DIR))
-CARD_SEARCH_PATH = [
+_BUILTIN_CARD_DIRS = [
Path(RESOURCES_DIR) / "inference_model_cards",
Path(RESOURCES_DIR) / "image_model_cards",
- _custom_cards_dir,
]
_card_cache: dict[ModelId, "ModelCard"] = {}
-async def _refresh_card_cache():
- for path in CARD_SEARCH_PATH:
- async for toml_file in path.rglob("*.toml"):
- try:
- card = await ModelCard.load_from_path(toml_file)
- if card.model_id not in _card_cache:
- _card_cache[card.model_id] = card
- except (ValidationError, TOMLKitError):
- pass
+async def _load_cards_from_dir(directory: Path, *, is_custom: bool) -> None:
+ """Load all TOML model cards from a directory into the cache."""
+ async for toml_file in directory.rglob("*.toml"):
+ try:
+ card = await ModelCard.load_from_path(toml_file)
+ if is_custom:
+ card = card.model_copy(update={"is_custom": True})
+ if card.model_id not in _card_cache:
+ _card_cache[card.model_id] = card
+ except (ValidationError, TOMLKitError):
+ pass
+
+
+async def _refresh_card_cache() -> None:
+ for path in _BUILTIN_CARD_DIRS:
+ await _load_cards_from_dir(path, is_custom=False)
+ await _load_cards_from_dir(_custom_cards_dir, is_custom=True)
def _is_image_card(card: "ModelCard") -> bool:
return any(t in (ModelTask.TextToImage, ModelTask.ImageToImage) for t in card.tasks)
+def get_card(model_id: ModelId) -> "ModelCard | None":
+ """Look up a single model card from the cache by ID."""
+ return _card_cache.get(model_id)
+
+
async def get_model_cards() -> list["ModelCard"]:
if len(_card_cache) == 0:
await _refresh_card_cache()
@@ -92,6 +104,7 @@ class ModelCard(CamelCaseModel):
capabilities: list[str] = []
uses_cfg: bool = False
trust_remote_code: bool = True
+ is_custom: bool = False
@field_validator("tasks", mode="before")
@classmethod
@@ -100,7 +113,7 @@ class ModelCard(CamelCaseModel):
async def save(self, path: Path) -> None:
async with await open_file(path, "w") as f:
- py = self.model_dump(exclude_none=True)
+ py = self.model_dump(exclude_none=True, exclude={"is_custom"})
data = tomlkit.dumps(py) # pyright: ignore[reportUnknownMemberType]
await f.write(data)
@@ -122,17 +135,24 @@ class ModelCard(CamelCaseModel):
if (mc := _card_cache.get(model_id)) is not None:
return mc
- return await ModelCard.fetch_from_hf(model_id)
+ mc = await ModelCard.fetch_from_hf(model_id)
+ await mc.save_to_custom_dir()
+ _card_cache[model_id] = mc
+ return mc
@staticmethod
async def fetch_from_hf(model_id: ModelId) -> "ModelCard":
- """Fetches storage size and number of layers for a Hugging Face model, returns Pydantic ModelMeta."""
+ """Fetches storage size and number of layers for a Hugging Face model, returns Pydantic ModelMeta.
+
+ This is a pure fetch — it does NOT save to disk or update the cache.
+ Persistence is handled by the event-sourcing layer (worker event handler).
+ """
# TODO: failure if files do not exist
config_data = await fetch_config_data(model_id)
num_layers = config_data.layer_count
mem_size_bytes = await fetch_safetensors_size(model_id)
- mc = ModelCard(
+ return ModelCard(
model_id=ModelId(model_id),
storage_size=mem_size_bytes,
n_layers=num_layers,
@@ -141,10 +161,13 @@ class ModelCard(CamelCaseModel):
num_key_value_heads=config_data.num_key_value_heads,
tasks=[ModelTask.TextGeneration],
trust_remote_code=False,
+ is_custom=True,
)
- await mc.save_to_custom_dir()
- _card_cache[model_id] = mc
- return mc
+
+
+def add_to_card_cache(card: "ModelCard") -> None:
+ """Add or update a model card in the in-memory cache."""
+ _card_cache[card.model_id] = card
async def delete_custom_card(model_id: ModelId) -> bool:
@@ -157,16 +180,6 @@ async def delete_custom_card(model_id: ModelId) -> bool:
return False
-def is_custom_card(model_id: ModelId) -> bool:
- """Check if a model card exists in the custom cards directory."""
- import os
-
- card_path = Path(str(EXO_CUSTOM_MODEL_CARDS_DIR)) / (
- ModelId(model_id).normalize() + ".toml"
- )
- return os.path.isfile(str(card_path))
-
-
class ConfigData(BaseModel):
model_config = {"extra": "ignore"} # Allow unknown fields
diff --git a/src/exo/shared/types/commands.py b/src/exo/shared/types/commands.py
index fb55632e..a6c988a5 100644
--- a/src/exo/shared/types/commands.py
+++ b/src/exo/shared/types/commands.py
@@ -81,6 +81,14 @@ class CancelDownload(BaseCommand):
model_id: ModelId
+class AddCustomModelCard(BaseCommand):
+ model_card: ModelCard
+
+
+class DeleteCustomModelCard(BaseCommand):
+ model_id: ModelId
+
+
DownloadCommand = StartDownload | DeleteDownload | CancelDownload
@@ -96,6 +104,8 @@ Command = (
| TaskCancelled
| TaskFinished
| SendInputChunk
+ | AddCustomModelCard
+ | DeleteCustomModelCard
)
diff --git a/src/exo/shared/types/events.py b/src/exo/shared/types/events.py
index c9c2fc3a..d9799313 100644
--- a/src/exo/shared/types/events.py
+++ b/src/exo/shared/types/events.py
@@ -3,9 +3,10 @@ from typing import final
from pydantic import Field
+from exo.shared.models.model_cards import ModelCard
from exo.shared.topology import Connection
from exo.shared.types.chunks import GenerationChunk, InputImageChunk
-from exo.shared.types.common import CommandId, Id, NodeId, SessionId, SystemId
+from exo.shared.types.common import CommandId, Id, ModelId, NodeId, SessionId, SystemId
from exo.shared.types.tasks import Task, TaskId, TaskStatus
from exo.shared.types.worker.downloads import DownloadProgress
from exo.shared.types.worker.instances import Instance, InstanceId
@@ -106,6 +107,14 @@ class TopologyEdgeDeleted(BaseEvent):
conn: Connection
+class CustomModelCardAdded(BaseEvent):
+ model_card: ModelCard
+
+
+class CustomModelCardDeleted(BaseEvent):
+ model_id: ModelId
+
+
@final
class TraceEventData(FrozenModel):
name: str
@@ -147,6 +156,8 @@ Event = (
| TopologyEdgeDeleted
| TracesCollected
| TracesMerged
+ | CustomModelCardAdded
+ | CustomModelCardDeleted
)
diff --git a/src/exo/worker/main.py b/src/exo/worker/main.py
index 5389e78d..4bfed825 100644
--- a/src/exo/worker/main.py
+++ b/src/exo/worker/main.py
@@ -8,7 +8,7 @@ from loguru import logger
from exo.api.types import ImageEditsTaskParams
from exo.download.download_utils import resolve_model_in_path
from exo.shared.apply import apply
-from exo.shared.models.model_cards import ModelId
+from exo.shared.models.model_cards import ModelId, add_to_card_cache, delete_custom_card
from exo.shared.types.commands import (
ForwarderCommand,
ForwarderDownloadCommand,
@@ -16,6 +16,8 @@ from exo.shared.types.commands import (
)
from exo.shared.types.common import CommandId, NodeId, SystemId
from exo.shared.types.events import (
+ CustomModelCardAdded,
+ CustomModelCardDeleted,
Event,
IndexedEvent,
InputChunkReceived,
@@ -130,6 +132,13 @@ class Worker:
event.chunk.data
)
+ if isinstance(event, CustomModelCardAdded):
+ await event.model_card.save_to_custom_dir()
+ add_to_card_cache(event.model_card)
+
+ if isinstance(event, CustomModelCardDeleted):
+ await delete_custom_card(event.model_id)
+
async def plan_step(self):
while True:
await anyio.sleep(0.1)
← e06e70a8 Prefer higher model download % for placement (#1767)
·
back to Exo
·
Fix first start in offline mode crash (#1782) 248919c2 →