[object Object]

← back to Exo

swap camelcasemodels for frozenmodels globally (#1957)

df332035ef377d8775f329f44e0b3b3d70b82ffb · 2026-04-22 12:49:25 +0100 · Evan Quiney

Files touched

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 →