← back to Exo
swap camelcasemodels for frozenmodels globally (#1957)
df332035ef377d8775f329f44e0b3b3d70b82ffb · 2026-04-22 12:49:25 +0100 · Evan Quiney
Files touched
M docs/architecture.mdM src/exo/api/types/api.pyM src/exo/main.pyM src/exo/master/main.pyM src/exo/master/placement.pyM src/exo/master/tests/test_placement.pyM src/exo/routing/connection_message.pyM src/exo/routing/router.pyM src/exo/routing/topics.pyM src/exo/shared/election.pyM src/exo/shared/models/model_cards.pyM src/exo/shared/types/commands.pyM src/exo/shared/types/common.pyM src/exo/shared/types/events.pyM src/exo/shared/types/profiling.pyM src/exo/shared/types/state.pyM src/exo/shared/types/thunderbolt.pyM src/exo/shared/types/worker/downloads.pyM src/exo/shared/types/worker/instances.pyM src/exo/shared/types/worker/runners.pyM src/exo/utils/pydantic_ext.pyM src/exo/worker/tests/unittests/test_runner/test_event_ordering.py
Diff
commit df332035ef377d8775f329f44e0b3b3d70b82ffb
Author: Evan Quiney <evanev7@gmail.com>
Date: Wed Apr 22 12:49:25 2026 +0100
swap camelcasemodels for frozenmodels globally (#1957)
---
docs/architecture.md | 2 +-
src/exo/api/types/api.py | 30 +++++++++++-----------
src/exo/main.py | 4 +--
src/exo/master/main.py | 8 ++++--
src/exo/master/placement.py | 8 ++++--
src/exo/master/tests/test_placement.py | 28 ++++++++++++--------
src/exo/routing/connection_message.py | 4 +--
src/exo/routing/router.py | 16 ++++++------
src/exo/routing/topics.py | 4 +--
src/exo/shared/election.py | 6 ++---
src/exo/shared/models/model_cards.py | 10 ++++----
src/exo/shared/types/commands.py | 6 ++---
src/exo/shared/types/common.py | 6 ++---
src/exo/shared/types/events.py | 8 +++---
src/exo/shared/types/profiling.py | 20 +++++++--------
src/exo/shared/types/state.py | 4 +--
src/exo/shared/types/thunderbolt.py | 6 ++---
src/exo/shared/types/worker/downloads.py | 4 +--
src/exo/shared/types/worker/instances.py | 4 +--
src/exo/shared/types/worker/runners.py | 4 +--
src/exo/utils/pydantic_ext.py | 27 +++++--------------
.../unittests/test_runner/test_event_ordering.py | 2 +-
22 files changed, 106 insertions(+), 105 deletions(-)
diff --git a/docs/architecture.md b/docs/architecture.md
index daf5d3f2..55983168 100644
--- a/docs/architecture.md
+++ b/docs/architecture.md
@@ -81,4 +81,4 @@ Whenever a device produces side effects, it captures those side effects in an `E
## Purity
-A significant goal of the current design is to make data flow explicit. Classes should either represent simple data (`CamelCaseModel`s typically, and `TaggedModel`s for unions) or active `System`s (Erlang `Actor`s), with all transformations of that data being "referentially transparent" - destructure and construct new data, don't mutate in place. We have had varying degrees of success with this, and are still exploring where purity makes sense.
+A significant goal of the current design is to make data flow explicit. Classes should either represent simple data (`FrozenModel`s typically, and `TaggedModel`s for unions) or active `System`s (Erlang `Actor`s), with all transformations of that data being "referentially transparent" - destructure and construct new data, don't mutate in place. We have had varying degrees of success with this, and are still exploring where purity makes sense.
diff --git a/src/exo/api/types/api.py b/src/exo/api/types/api.py
index 75b820c3..ddf07f78 100644
--- a/src/exo/api/types/api.py
+++ b/src/exo/api/types/api.py
@@ -11,7 +11,7 @@ from exo.shared.types.memory import Memory
from exo.shared.types.text_generation import ReasoningEffort
from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta
from exo.shared.types.worker.shards import Sharding, ShardMetadata
-from exo.utils.pydantic_ext import CamelCaseModel
+from exo.utils.pydantic_ext import FrozenModel
FinishReason = Literal[
"stop", "length", "tool_calls", "content_filter", "function_call", "error"
@@ -418,29 +418,29 @@ class ImageListResponse(BaseModel, frozen=True):
data: list[ImageListItem]
-class StartDownloadParams(CamelCaseModel):
+class StartDownloadParams(FrozenModel):
target_node_id: NodeId
shard_metadata: ShardMetadata
-class StartDownloadResponse(CamelCaseModel):
+class StartDownloadResponse(FrozenModel):
command_id: CommandId
-class DeleteDownloadResponse(CamelCaseModel):
+class DeleteDownloadResponse(FrozenModel):
command_id: CommandId
-class CancelDownloadParams(CamelCaseModel):
+class CancelDownloadParams(FrozenModel):
target_node_id: NodeId
model_id: ModelId
-class CancelDownloadResponse(CamelCaseModel):
+class CancelDownloadResponse(FrozenModel):
command_id: CommandId
-class TraceEventResponse(CamelCaseModel):
+class TraceEventResponse(FrozenModel):
name: str
start_us: int
duration_us: int
@@ -448,12 +448,12 @@ class TraceEventResponse(CamelCaseModel):
category: str
-class TraceResponse(CamelCaseModel):
+class TraceResponse(FrozenModel):
task_id: str
traces: list[TraceEventResponse]
-class TraceCategoryStats(CamelCaseModel):
+class TraceCategoryStats(FrozenModel):
total_us: int
count: int
min_us: int
@@ -461,31 +461,31 @@ class TraceCategoryStats(CamelCaseModel):
avg_us: float
-class TraceRankStats(CamelCaseModel):
+class TraceRankStats(FrozenModel):
by_category: dict[str, TraceCategoryStats]
-class TraceStatsResponse(CamelCaseModel):
+class TraceStatsResponse(FrozenModel):
task_id: str
total_wall_time_us: int
by_category: dict[str, TraceCategoryStats]
by_rank: dict[int, TraceRankStats]
-class TraceListItem(CamelCaseModel):
+class TraceListItem(FrozenModel):
task_id: str
created_at: str
file_size: int
-class TraceListResponse(CamelCaseModel):
+class TraceListResponse(FrozenModel):
traces: list[TraceListItem]
-class DeleteTracesRequest(CamelCaseModel):
+class DeleteTracesRequest(FrozenModel):
task_ids: list[str]
-class DeleteTracesResponse(CamelCaseModel):
+class DeleteTracesResponse(FrozenModel):
deleted: list[str]
not_found: list[str]
diff --git a/src/exo/main.py b/src/exo/main.py
index 903cca6f..30c54e29 100644
--- a/src/exo/main.py
+++ b/src/exo/main.py
@@ -22,7 +22,7 @@ from exo.shared.election import Election, ElectionResult
from exo.shared.logging import logger_cleanup, logger_setup
from exo.shared.types.common import NodeId, SessionId
from exo.utils.channels import Receiver, channel
-from exo.utils.pydantic_ext import CamelCaseModel
+from exo.utils.pydantic_ext import FrozenModel
from exo.utils.task_group import TaskGroup
from exo.worker.main import Worker
@@ -308,7 +308,7 @@ def main():
logger_cleanup()
-class Args(CamelCaseModel):
+class Args(FrozenModel):
verbosity: int = 0
force_master: bool = False
spawn_api: bool = False
diff --git a/src/exo/master/main.py b/src/exo/master/main.py
index 1511897e..0a2afa5a 100644
--- a/src/exo/master/main.py
+++ b/src/exo/master/main.py
@@ -413,9 +413,13 @@ class Master:
indexed = IndexedEvent(event=event, idx=len(self._event_log))
self.state = apply(self.state, indexed)
- event._master_time_stamp = datetime.now(tz=timezone.utc) # pyright: ignore[reportPrivateUsage]
+ event = event.model_copy(
+ update={"_master_time_stamp": datetime.now(tz=timezone.utc)}
+ )
if isinstance(event, NodeGatheredInfo):
- event.when = str(datetime.now(tz=timezone.utc))
+ event = event.model_copy(
+ update={"when": str(datetime.now(tz=timezone.utc))}
+ )
self._event_log.append(event)
await self._send_event(indexed)
diff --git a/src/exo/master/placement.py b/src/exo/master/placement.py
index a8eae188..160a010f 100644
--- a/src/exo/master/placement.py
+++ b/src/exo/master/placement.py
@@ -204,8 +204,12 @@ def place_instance(
# Single-node: force Pipeline/Ring (Tensor and Jaccl require multi-node)
if len(selected_cycle) == 1:
- command.instance_meta = InstanceMeta.MlxRing
- command.sharding = Sharding.Pipeline
+ command = command.model_copy(
+ update={
+ "instance_meta": InstanceMeta.MlxRing,
+ "sharding": Sharding.Pipeline,
+ }
+ )
shard_assignments = get_shard_assignments(
command.model_card, selected_cycle, command.sharding, node_memory
diff --git a/src/exo/master/tests/test_placement.py b/src/exo/master/tests/test_placement.py
index 40530ad2..3e6b9d92 100644
--- a/src/exo/master/tests/test_placement.py
+++ b/src/exo/master/tests/test_placement.py
@@ -95,10 +95,14 @@ def test_get_instance_placements_create_instance(
model_card: ModelCard,
):
# arrange
- model_card.n_layers = total_layers
- model_card.storage_size = Memory.from_bytes(
- sum(available_memory)
- ) # make it exactly fit across all nodes
+ model_card = model_card.model_copy(
+ update={
+ "n_layers": total_layers,
+ "storage_size": Memory.from_bytes(
+ sum(available_memory)
+ ), # make it exactly fit across all nodes
+ }
+ )
topology = Topology()
cic = place_instance_command(model_card)
@@ -296,7 +300,7 @@ def test_placement_selects_leaf_nodes(
# arrange
topology = Topology()
- model_card.storage_size = Memory.from_bytes(1000)
+ model_card = model_card.model_copy(update={"storage_size": Memory.from_bytes(1000)})
node_id_a = NodeId()
node_id_b = NodeId()
@@ -364,8 +368,12 @@ def test_tensor_rdma_backend_connectivity_matrix(
):
# arrange
topology = Topology()
- model_card.n_layers = 12
- model_card.storage_size = Memory.from_bytes(1500)
+ model_card = model_card.model_copy(
+ update={
+ "n_layers": 12,
+ "storage_size": Memory.from_bytes(1500),
+ }
+ )
node_a = NodeId()
node_b = NodeId()
@@ -605,7 +613,7 @@ def test_placement_prefers_cycle_with_downloaded_model(
"""When two cycles are otherwise equal, prefer the one with the model already downloaded."""
topology = Topology()
- model_card.storage_size = Memory.from_bytes(500)
+ model_card = model_card.model_copy(update={"storage_size": Memory.from_bytes(500)})
node_a = NodeId()
node_b = NodeId()
@@ -653,7 +661,7 @@ def test_placement_prefers_cycle_with_higher_download_progress(
"""When two cycles are otherwise equal, prefer the one with more download progress."""
topology = Topology()
- model_card.storage_size = Memory.from_bytes(1000)
+ model_card = model_card.model_copy(update={"storage_size": Memory.from_bytes(1000)})
node_a = NodeId()
node_b = NodeId()
@@ -725,7 +733,7 @@ def test_placement_does_not_prefer_cycle_with_failed_download(
"""A failed download should count as 0% — not preferred over a node with no download history."""
topology = Topology()
- model_card.storage_size = Memory.from_bytes(500)
+ model_card = model_card.model_copy(update={"storage_size": Memory.from_bytes(500)})
node_a = NodeId()
node_b = NodeId()
diff --git a/src/exo/routing/connection_message.py b/src/exo/routing/connection_message.py
index 3cc0362d..b0089112 100644
--- a/src/exo/routing/connection_message.py
+++ b/src/exo/routing/connection_message.py
@@ -1,12 +1,12 @@
from exo_pyo3_bindings import PyFromSwarm
from exo.shared.types.common import NodeId
-from exo.utils.pydantic_ext import CamelCaseModel
+from exo.utils.pydantic_ext import FrozenModel
"""Serialisable types for Connection Updates/Messages"""
-class ConnectionMessage(CamelCaseModel):
+class ConnectionMessage(FrozenModel):
node_id: NodeId
connected: bool
diff --git a/src/exo/routing/router.py b/src/exo/routing/router.py
index 9447d6aa..a9341d10 100644
--- a/src/exo/routing/router.py
+++ b/src/exo/routing/router.py
@@ -25,7 +25,7 @@ from loguru import logger
from exo.shared.constants import EXO_NODE_ID_KEYPAIR
from exo.utils.channels import Receiver, Sender, channel
-from exo.utils.pydantic_ext import CamelCaseModel
+from exo.utils.pydantic_ext import FrozenModel
from exo.utils.task_group import TaskGroup
from .connection_message import ConnectionMessage
@@ -36,7 +36,7 @@ from .topics import CONNECTION_MESSAGES, PublishPolicy, TypedTopic
# of preventing feedback, as it does not ask for a system id so cannot tell
# which message is coming/going to which system.
# This is currently only relevant for Election
-class TopicRouter[T: CamelCaseModel]:
+class TopicRouter[T: FrozenModel]:
def __init__(
self,
topic: TypedTopic[T],
@@ -114,7 +114,7 @@ class Router:
)
def __init__(self, handle: NetworkingHandle):
- self.topic_routers: dict[str, TopicRouter[CamelCaseModel]] = {}
+ self.topic_routers: dict[str, TopicRouter[FrozenModel]] = {}
send, recv = channel[tuple[str, bytes]]()
self.networking_receiver: Receiver[tuple[str, bytes]] = recv
self._net: NetworkingHandle = handle
@@ -122,18 +122,18 @@ class Router:
self._id_count = count()
self._tg: TaskGroup = TaskGroup()
- async def register_topic[T: CamelCaseModel](self, topic: TypedTopic[T]):
+ async def register_topic[T: FrozenModel](self, topic: TypedTopic[T]):
send = self._tmp_networking_sender
if send:
self._tmp_networking_sender = None
else:
send = self.networking_receiver.clone_sender()
router = TopicRouter[T](topic, send)
- self.topic_routers[topic.topic] = cast(TopicRouter[CamelCaseModel], router)
+ self.topic_routers[topic.topic] = cast(TopicRouter[FrozenModel], router)
if self._tg.is_running():
await self._networking_subscribe(topic.topic)
- def sender[T: CamelCaseModel](self, topic: TypedTopic[T]) -> Sender[T]:
+ def sender[T: FrozenModel](self, topic: TypedTopic[T]) -> Sender[T]:
router = self.topic_routers.get(topic.topic, None)
# There's gotta be a way to do this without THIS many asserts
assert router is not None
@@ -141,7 +141,7 @@ class Router:
sender = cast(TopicRouter[T], router).new_sender()
return sender
- def receiver[T: CamelCaseModel](self, topic: TypedTopic[T]) -> Receiver[T]:
+ def receiver[T: FrozenModel](self, topic: TypedTopic[T]) -> Receiver[T]:
router = self.topic_routers.get(topic.topic, None)
# There's gotta be a way to do this without THIS many asserts
@@ -150,7 +150,7 @@ class Router:
assert router.topic.model_type == topic.model_type
send, recv = channel[T]()
- router.senders.add(cast(Sender[CamelCaseModel], send))
+ router.senders.add(cast(Sender[FrozenModel], send))
return recv
diff --git a/src/exo/routing/topics.py b/src/exo/routing/topics.py
index 5d95a2a1..9776e542 100644
--- a/src/exo/routing/topics.py
+++ b/src/exo/routing/topics.py
@@ -8,7 +8,7 @@ from exo.shared.types.events import (
GlobalForwarderEvent,
LocalForwarderEvent,
)
-from exo.utils.pydantic_ext import CamelCaseModel
+from exo.utils.pydantic_ext import FrozenModel
class PublishPolicy(str, Enum):
@@ -21,7 +21,7 @@ class PublishPolicy(str, Enum):
@dataclass # (frozen=True)
-class TypedTopic[T: CamelCaseModel]:
+class TypedTopic[T: FrozenModel]:
topic: str
publish_policy: PublishPolicy
diff --git a/src/exo/shared/election.py b/src/exo/shared/election.py
index 6f6e1f8f..958a83d2 100644
--- a/src/exo/shared/election.py
+++ b/src/exo/shared/election.py
@@ -12,13 +12,13 @@ from exo.routing.connection_message import ConnectionMessage
from exo.shared.types.commands import ForwarderCommand
from exo.shared.types.common import NodeId, SessionId
from exo.utils.channels import Receiver, Sender
-from exo.utils.pydantic_ext import CamelCaseModel
+from exo.utils.pydantic_ext import FrozenModel
from exo.utils.task_group import TaskGroup
DEFAULT_ELECTION_TIMEOUT = 3.0
-class ElectionMessage(CamelCaseModel):
+class ElectionMessage(FrozenModel):
clock: int
seniority: int
proposed_session: SessionId
@@ -39,7 +39,7 @@ class ElectionMessage(CamelCaseModel):
)
-class ElectionResult(CamelCaseModel):
+class ElectionResult(FrozenModel):
session_id: SessionId
won_clock: int
is_new_master: bool
diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py
index 7e9f9c30..d5a8e172 100644
--- a/src/exo/shared/models/model_cards.py
+++ b/src/exo/shared/models/model_cards.py
@@ -28,7 +28,7 @@ from exo.shared.constants import (
)
from exo.shared.types.common import ModelId
from exo.shared.types.memory import Memory
-from exo.utils.pydantic_ext import CamelCaseModel
+from exo.utils.pydantic_ext import FrozenModel
# kinda ugly...
# TODO: load search path from config.toml
@@ -100,7 +100,7 @@ class ModelTask(str, Enum):
ImageToImage = "ImageToImage"
-class ComponentInfo(CamelCaseModel):
+class ComponentInfo(FrozenModel):
component_name: str
component_path: str
storage_size: Memory
@@ -109,7 +109,7 @@ class ComponentInfo(CamelCaseModel):
safetensors_index_filename: str | None = None
-class VisionCardConfig(CamelCaseModel):
+class VisionCardConfig(FrozenModel):
image_token_id: int
model_type: str
weights_repo: str = ""
@@ -117,7 +117,7 @@ class VisionCardConfig(CamelCaseModel):
processor_repo: str | None = None
-class SamplingValues(CamelCaseModel):
+class SamplingValues(FrozenModel):
temperature: float | None = None
top_p: float | None = None
top_k: int | None = None
@@ -132,7 +132,7 @@ class SamplingDefaults(SamplingValues):
non_thinking: SamplingValues | None = None
-class ModelCard(CamelCaseModel):
+class ModelCard(FrozenModel):
model_id: ModelId
storage_size: Memory
n_layers: PositiveInt
diff --git a/src/exo/shared/types/commands.py b/src/exo/shared/types/commands.py
index a6c988a5..b2dc3c89 100644
--- a/src/exo/shared/types/commands.py
+++ b/src/exo/shared/types/commands.py
@@ -10,7 +10,7 @@ from exo.shared.types.common import CommandId, NodeId, SystemId
from exo.shared.types.text_generation import TextGenerationTaskParams
from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta
from exo.shared.types.worker.shards import Sharding, ShardMetadata
-from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
+from exo.utils.pydantic_ext import FrozenModel, TaggedModel
class BaseCommand(TaggedModel):
@@ -109,11 +109,11 @@ Command = (
)
-class ForwarderCommand(CamelCaseModel):
+class ForwarderCommand(FrozenModel):
origin: SystemId
command: Command
-class ForwarderDownloadCommand(CamelCaseModel):
+class ForwarderDownloadCommand(FrozenModel):
origin: SystemId
command: DownloadCommand
diff --git a/src/exo/shared/types/common.py b/src/exo/shared/types/common.py
index e539386f..097803d3 100644
--- a/src/exo/shared/types/common.py
+++ b/src/exo/shared/types/common.py
@@ -4,7 +4,7 @@ from uuid import uuid4
from pydantic import GetCoreSchemaHandler, field_validator
from pydantic_core import CoreSchema, core_schema
-from exo.utils.pydantic_ext import CamelCaseModel
+from exo.utils.pydantic_ext import FrozenModel
class Id(str):
@@ -59,12 +59,12 @@ class TruncatingString(str):
)
-class SessionId(CamelCaseModel):
+class SessionId(FrozenModel):
master_node_id: NodeId
election_clock: int
-class Host(CamelCaseModel):
+class Host(FrozenModel):
ip: str
port: int
diff --git a/src/exo/shared/types/events.py b/src/exo/shared/types/events.py
index d9799313..b750a2ae 100644
--- a/src/exo/shared/types/events.py
+++ b/src/exo/shared/types/events.py
@@ -12,7 +12,7 @@ from exo.shared.types.worker.downloads import DownloadProgress
from exo.shared.types.worker.instances import Instance, InstanceId
from exo.shared.types.worker.runners import RunnerId, RunnerStatus
from exo.utils.info_gatherer.info_gatherer import GatheredInfo
-from exo.utils.pydantic_ext import CamelCaseModel, FrozenModel, TaggedModel
+from exo.utils.pydantic_ext import FrozenModel, TaggedModel
class EventId(Id):
@@ -161,14 +161,14 @@ Event = (
)
-class IndexedEvent(CamelCaseModel):
+class IndexedEvent(FrozenModel):
"""An event indexed by the master, with a globally unique index"""
idx: int = Field(ge=0)
event: Event
-class GlobalForwarderEvent(CamelCaseModel):
+class GlobalForwarderEvent(FrozenModel):
"""An event the forwarder will serialize and send over the network"""
origin_idx: int = Field(ge=0)
@@ -177,7 +177,7 @@ class GlobalForwarderEvent(CamelCaseModel):
event: Event
-class LocalForwarderEvent(CamelCaseModel):
+class LocalForwarderEvent(FrozenModel):
"""An event the forwarder will serialize and send over the network"""
origin_idx: int = Field(ge=0)
diff --git a/src/exo/shared/types/profiling.py b/src/exo/shared/types/profiling.py
index ad1d48c0..a3548dc9 100644
--- a/src/exo/shared/types/profiling.py
+++ b/src/exo/shared/types/profiling.py
@@ -7,10 +7,10 @@ import psutil
from exo.shared.types.memory import Memory
from exo.shared.types.thunderbolt import ThunderboltIdentifier
-from exo.utils.pydantic_ext import CamelCaseModel
+from exo.utils.pydantic_ext import FrozenModel
-class MemoryUsage(CamelCaseModel):
+class MemoryUsage(FrozenModel):
ram_total: Memory
ram_available: Memory
swap_total: Memory
@@ -40,7 +40,7 @@ class MemoryUsage(CamelCaseModel):
)
-class DiskUsage(CamelCaseModel):
+class DiskUsage(FrozenModel):
"""Disk space usage for the models directory."""
total: Memory
@@ -56,7 +56,7 @@ class DiskUsage(CamelCaseModel):
)
-class SystemPerformanceProfile(CamelCaseModel):
+class SystemPerformanceProfile(FrozenModel):
# TODO: flops_fp16: float
gpu_usage: float = 0.0
@@ -69,13 +69,13 @@ class SystemPerformanceProfile(CamelCaseModel):
InterfaceType = Literal["wifi", "ethernet", "maybe_ethernet", "thunderbolt", "unknown"]
-class NetworkInterfaceInfo(CamelCaseModel):
+class NetworkInterfaceInfo(FrozenModel):
name: str
ip_address: str
interface_type: InterfaceType = "unknown"
-class NodeIdentity(CamelCaseModel):
+class NodeIdentity(FrozenModel):
"""Static and slow-changing node identification data."""
model_id: str = "Unknown"
@@ -85,25 +85,25 @@ class NodeIdentity(CamelCaseModel):
os_build_version: str = "Unknown"
-class NodeNetworkInfo(CamelCaseModel):
+class NodeNetworkInfo(FrozenModel):
"""Network interface information for a node."""
interfaces: Sequence[NetworkInterfaceInfo] = []
-class NodeThunderboltInfo(CamelCaseModel):
+class NodeThunderboltInfo(FrozenModel):
"""Thunderbolt interface identifiers for a node."""
interfaces: Sequence[ThunderboltIdentifier] = []
-class NodeRdmaCtlStatus(CamelCaseModel):
+class NodeRdmaCtlStatus(FrozenModel):
"""Whether RDMA is enabled on this node (via rdma_ctl)."""
enabled: bool
-class ThunderboltBridgeStatus(CamelCaseModel):
+class ThunderboltBridgeStatus(FrozenModel):
"""Whether the Thunderbolt Bridge network service is enabled on this node."""
enabled: bool
diff --git a/src/exo/shared/types/state.py b/src/exo/shared/types/state.py
index 7350cfb0..71cffae7 100644
--- a/src/exo/shared/types/state.py
+++ b/src/exo/shared/types/state.py
@@ -21,10 +21,10 @@ from exo.shared.types.tasks import Task, TaskId
from exo.shared.types.worker.downloads import DownloadProgress
from exo.shared.types.worker.instances import Instance, InstanceId
from exo.shared.types.worker.runners import RunnerId, RunnerStatus
-from exo.utils.pydantic_ext import CamelCaseModel
+from exo.utils.pydantic_ext import FrozenModel
-class State(CamelCaseModel):
+class State(FrozenModel):
"""Global system state.
The :class:`Topology` instance is encoded/decoded via an immutable
diff --git a/src/exo/shared/types/thunderbolt.py b/src/exo/shared/types/thunderbolt.py
index 809cf9fa..34cd1cca 100644
--- a/src/exo/shared/types/thunderbolt.py
+++ b/src/exo/shared/types/thunderbolt.py
@@ -1,15 +1,15 @@
import anyio
from pydantic import BaseModel, Field
-from exo.utils.pydantic_ext import CamelCaseModel
+from exo.utils.pydantic_ext import FrozenModel
-class ThunderboltConnection(CamelCaseModel):
+class ThunderboltConnection(FrozenModel):
source_uuid: str
sink_uuid: str
-class ThunderboltIdentifier(CamelCaseModel):
+class ThunderboltIdentifier(FrozenModel):
rdma_interface: str
domain_uuid: str
link_speed: str = ""
diff --git a/src/exo/shared/types/worker/downloads.py b/src/exo/shared/types/worker/downloads.py
index 52036c0b..938be82e 100644
--- a/src/exo/shared/types/worker/downloads.py
+++ b/src/exo/shared/types/worker/downloads.py
@@ -6,10 +6,10 @@ from pydantic import BaseModel, ConfigDict, Field, PositiveInt
from exo.shared.types.common import NodeId
from exo.shared.types.memory import Memory
from exo.shared.types.worker.shards import ShardMetadata
-from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
+from exo.utils.pydantic_ext import FrozenModel, TaggedModel
-class DownloadProgressData(CamelCaseModel):
+class DownloadProgressData(FrozenModel):
total: Memory
downloaded: Memory
downloaded_this_session: Memory
diff --git a/src/exo/shared/types/worker/instances.py b/src/exo/shared/types/worker/instances.py
index 76bd6fd4..16233f3f 100644
--- a/src/exo/shared/types/worker/instances.py
+++ b/src/exo/shared/types/worker/instances.py
@@ -5,7 +5,7 @@ from pydantic import model_validator
from exo.shared.models.model_cards import ModelTask
from exo.shared.types.common import Host, Id, NodeId
from exo.shared.types.worker.runners import RunnerId, ShardAssignments, ShardMetadata
-from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
+from exo.utils.pydantic_ext import FrozenModel, TaggedModel
class InstanceId(Id):
@@ -39,7 +39,7 @@ class MlxJacclInstance(BaseInstance):
Instance = MlxRingInstance | MlxJacclInstance
-class BoundInstance(CamelCaseModel):
+class BoundInstance(FrozenModel):
instance: Instance
bound_runner_id: RunnerId
bound_node_id: NodeId
diff --git a/src/exo/shared/types/worker/runners.py b/src/exo/shared/types/worker/runners.py
index 1ac68947..4875c6e5 100644
--- a/src/exo/shared/types/worker/runners.py
+++ b/src/exo/shared/types/worker/runners.py
@@ -5,7 +5,7 @@ from pydantic import model_validator
from exo.shared.models.model_cards import ModelId
from exo.shared.types.common import Id, NodeId
from exo.shared.types.worker.shards import ShardMetadata
-from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
+from exo.utils.pydantic_ext import FrozenModel, TaggedModel
class RunnerId(Id):
@@ -81,7 +81,7 @@ RunnerStatus = (
)
-class ShardAssignments(CamelCaseModel):
+class ShardAssignments(FrozenModel):
model_id: ModelId
runner_to_shard: Mapping[RunnerId, ShardMetadata]
node_to_runner: Mapping[NodeId, RunnerId]
diff --git a/src/exo/utils/pydantic_ext.py b/src/exo/utils/pydantic_ext.py
index 07c8dc5e..e8c6068e 100644
--- a/src/exo/utils/pydantic_ext.py
+++ b/src/exo/utils/pydantic_ext.py
@@ -1,5 +1,3 @@
-# pyright: reportAny=false, reportUnknownArgumentType=false, reportUnknownVariableType=false
-
from typing import Any, Self
from pydantic import BaseModel, ConfigDict, model_serializer, model_validator
@@ -10,19 +8,6 @@ from pydantic_core.core_schema import (
)
-class CamelCaseModel(BaseModel):
- """
- A model whose fields are aliased to camel-case from snake-case.
- """
-
- model_config = ConfigDict(
- alias_generator=to_camel,
- validate_by_name=True,
- extra="forbid",
- strict=True,
- )
-
-
class FrozenModel(BaseModel):
model_config = ConfigDict(
alias_generator=to_camel,
@@ -33,19 +18,19 @@ class FrozenModel(BaseModel):
)
-class TaggedModel(CamelCaseModel):
+class TaggedModel(FrozenModel):
@model_serializer(mode="wrap")
def _serialize(self, handler: SerializerFunctionWrapHandler):
- inner = handler(self)
+ inner = handler(self) # pyright: ignore[reportAny]
return {self.__class__.__name__: inner}
@model_validator(mode="wrap")
@classmethod
- def _validate(cls, v: Any, handler: ValidatorFunctionWrapHandler) -> Self:
- if isinstance(v, dict) and len(v) == 1 and cls.__name__ in v:
- return handler(v[cls.__name__])
+ def _validate(cls, v: Any, handler: ValidatorFunctionWrapHandler) -> Self: # pyright: ignore[reportAny]
+ if isinstance(v, dict) and len(v) == 1 and cls.__name__ in v: # pyright: ignore[reportUnknownArgumentType]
+ return handler(v[cls.__name__]) # pyright: ignore[reportAny]
- return handler(v)
+ return handler(v) # pyright: ignore[reportAny]
def __str__(self) -> str:
return f"{self.__class__.__name__}({super().__str__()})"
diff --git a/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py b/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py
index ffd8fbfd..62b0840e 100644
--- a/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py
+++ b/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py
@@ -111,7 +111,7 @@ CHAT_TASK = TextGeneration(
def assert_events_equal(test_events: Iterable[Event], true_events: Iterable[Event]):
for test_event, true_event in zip(test_events, true_events, strict=True):
- test_event.event_id = true_event.event_id
+ test_event = test_event.model_copy(update={"event_id": true_event.event_id})
assert test_event == true_event, f"{test_event} != {true_event}"
← af673845 Ignore HF remote repo changes (temporary fix) (#1958)
·
back to Exo
·
remove layer loading callback (#1890) 0a549f88 →