Add explicit MetaInstance binding, slim MetaInstance to use ModelId

- Add MetaInstanceBound event and meta_instance_backing State field
  for explicit MetaInstance → Instance binding (prevents ambiguous
  linking when two MetaInstances have identical constraints)
- Replace model_card: ModelCard with model_id: ModelId on MetaInstance
  (load ModelCard on-demand at placement time)
- Add MetaInstance API endpoints (POST /meta_instance, DELETE)
- Update dashboard to use MetaInstances as primary primitive with
  unified display items merging MetaInstances and orphan instances
- Dashboard launches via MetaInstance instead of direct Instance creation

Co-Authored-By: Claude Opus 4.6 <[email protected]>
This commit is contained in:
Alex Cheema
2026-02-10 15:53:07 -08:00
co-authored by Claude Opus 4.6
parent ceb76b8f6c
commit f4329c72c2
11 changed files with 550 additions and 77 deletions
+20
View File
@@ -232,6 +232,19 @@ interface RawStateResponse {
>;
// Thunderbolt bridge cycles (nodes with bridge enabled forming loops)
thunderboltBridgeCycles?: string[][];
// MetaInstances (declarative instance constraints)
metaInstances?: Record<string, MetaInstanceData>;
// Explicit MetaInstance → Instance binding
metaInstanceBacking?: Record<string, string>;
}
export interface MetaInstanceData {
metaInstanceId: string;
modelId: string;
sharding: string;
instanceMeta: string;
minNodes: number;
nodeIds: string[] | null;
}
export interface MessageAttachment {
@@ -495,6 +508,8 @@ class AppStore {
isLoadingPreviews = $state(false);
previewNodeFilter = $state<Set<string>>(new Set());
lastUpdate = $state<number | null>(null);
metaInstances = $state<Record<string, MetaInstanceData>>({});
metaInstanceBacking = $state<Record<string, string>>({});
thunderboltBridgeCycles = $state<string[][]>([]);
nodeThunderboltBridge = $state<
Record<
@@ -1206,6 +1221,9 @@ class AppStore {
if (data.downloads) {
this.downloads = data.downloads;
}
// MetaInstances
this.metaInstances = data.metaInstances ?? {};
this.metaInstanceBacking = data.metaInstanceBacking ?? {};
// Thunderbolt bridge cycles
this.thunderboltBridgeCycles = data.thunderboltBridgeCycles ?? [];
// Thunderbolt bridge status per node
@@ -2956,6 +2974,8 @@ export const tps = () => appStore.tps;
export const totalTokens = () => appStore.totalTokens;
export const topologyData = () => appStore.topologyData;
export const instances = () => appStore.instances;
export const metaInstances = () => appStore.metaInstances;
export const metaInstanceBacking = () => appStore.metaInstanceBacking;
export const runners = () => appStore.runners;
export const downloads = () => appStore.downloads;
export const placementPreviews = () => appStore.placementPreviews;
+186 -53
View File
@@ -37,10 +37,13 @@
toggleTopologyOnlyMode,
chatSidebarVisible,
toggleChatSidebarVisible,
metaInstances,
metaInstanceBacking,
thunderboltBridgeCycles,
nodeThunderboltBridge,
type DownloadProgress,
type PlacementPreview,
type MetaInstanceData,
} from "$lib/stores/app.svelte";
import HeaderNav from "$lib/components/HeaderNav.svelte";
import { fade, fly } from "svelte/transition";
@@ -60,6 +63,8 @@
const debugEnabled = $derived(debugMode());
const topologyOnlyEnabled = $derived(topologyOnlyMode());
const sidebarVisible = $derived(chatSidebarVisible());
const metaInstancesData = $derived(metaInstances());
const metaInstanceBackingData = $derived(metaInstanceBacking());
const tbBridgeCycles = $derived(thunderboltBridgeCycles());
const tbBridgeData = $derived(nodeThunderboltBridge());
const nodeFilter = $derived(previewNodeFilter());
@@ -593,39 +598,22 @@
launchingModelId = modelId;
try {
// Use the specific preview if provided, otherwise fall back to filtered preview
const preview = specificPreview ?? filteredPreview();
let instanceData: unknown;
if (preview?.instance) {
// Use the instance from the preview
instanceData = preview.instance;
} else {
// Fallback: GET placement from API
const placementResponse = await fetch(
`/instance/placement?model_id=${encodeURIComponent(modelId)}&sharding=${selectedSharding}&instance_meta=${selectedInstanceType}&min_nodes=${selectedMinNodes}`,
);
if (!placementResponse.ok) {
const errorText = await placementResponse.text();
console.error("Failed to get placement:", errorText);
return;
}
instanceData = await placementResponse.json();
}
// POST the instance to create it
const response = await fetch("/instance", {
const response = await fetch("/meta_instance", {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify({ instance: instanceData }),
body: JSON.stringify({
model_id: modelId,
sharding: preview?.sharding ?? selectedSharding,
instance_meta: preview?.instance_meta ?? selectedInstanceType,
min_nodes: selectedMinNodes,
}),
});
if (!response.ok) {
const errorText = await response.text();
console.error("Failed to launch instance:", errorText);
console.error("Failed to create meta instance:", errorText);
} else {
// Always auto-select the newly launched model so the user chats to what they just launched
setSelectedChatModel(modelId);
@@ -645,7 +633,7 @@
setTimeout(scrollToBottom, 1000);
}
} catch (error) {
console.error("Error launching instance:", error);
console.error("Error creating meta instance:", error);
} finally {
launchingModelId = null;
}
@@ -1148,6 +1136,64 @@
}
}
async function deleteMetaInstance(metaInstanceId: string) {
const meta = metaInstancesData[metaInstanceId];
const modelId = meta?.modelId ?? "unknown";
if (!confirm(`Delete model ${modelId}?`)) return;
const wasSelected = selectedChatModel() === modelId;
try {
const response = await fetch(`/meta_instance/${metaInstanceId}`, {
method: "DELETE",
headers: { "Content-Type": "application/json" },
});
if (!response.ok) {
console.error("Failed to delete meta instance:", response.status);
} else if (wasSelected) {
// Switch to another available model or clear selection
const remainingInstances = Object.entries(instanceData).filter(
([id]) => id !== getBackingInstanceId(metaInstanceId),
);
if (remainingInstances.length > 0) {
const [, lastInstance] =
remainingInstances[remainingInstances.length - 1];
const newModelId = getInstanceModelId(lastInstance);
if (
newModelId &&
newModelId !== "Unknown" &&
newModelId !== "Unknown Model"
) {
setSelectedChatModel(newModelId);
} else {
setSelectedChatModel("");
}
} else {
setSelectedChatModel("");
}
}
} catch (error) {
console.error("Error deleting meta instance:", error);
}
}
// Find the backing Instance ID for a MetaInstance via the explicit binding map
function getBackingInstanceId(metaInstanceId: string): string | null {
return metaInstanceBackingData[metaInstanceId] ?? null;
}
// Get the set of Instance IDs that are backing MetaInstances
function getBackedInstanceIds(): Set<string> {
return new Set(Object.values(metaInstanceBackingData));
}
// Get orphan Instance IDs (not backing any MetaInstance)
function getOrphanInstanceIds(): string[] {
const backedIds = getBackedInstanceIds();
return Object.keys(instanceData).filter((id) => !backedIds.has(id));
}
// Helper to unwrap tagged unions like { MlxRingInstance: {...} }
function getTagged(obj: unknown): [string | null, unknown] {
if (!obj || typeof obj !== "object") return [null, null];
@@ -1540,7 +1586,45 @@
}
const nodeCount = $derived(data ? Object.keys(data.nodes).length : 0);
const instanceCount = $derived(Object.keys(instanceData).length);
const metaInstanceCount = $derived(Object.keys(metaInstancesData).length);
const orphanInstanceIds = $derived(getOrphanInstanceIds());
const instanceCount = $derived(metaInstanceCount + orphanInstanceIds.length);
// Unified display items: MetaInstances first, then orphan Instances
interface DisplayItem {
id: string; // MetaInstance ID or Instance ID (used as key and displayed)
modelId: string;
instance: unknown | null; // The backing/orphan instance (tagged union) or null if placing
instanceId: string | null; // The actual Instance ID (for topology hover)
isMetaInstance: boolean;
}
const unifiedDisplayItems = $derived((): DisplayItem[] => {
const items: DisplayItem[] = [];
// MetaInstances
for (const [metaId, meta] of Object.entries(metaInstancesData)) {
const backingId = getBackingInstanceId(metaId);
items.push({
id: metaId,
modelId: meta.modelId,
instance: backingId ? instanceData[backingId] : null,
instanceId: backingId,
isMetaInstance: true,
});
}
// Orphan Instances
for (const orphanId of getOrphanInstanceIds()) {
const inst = instanceData[orphanId];
items.push({
id: orphanId,
modelId: getInstanceModelId(inst),
instance: inst,
instanceId: orphanId,
isMetaInstance: false,
});
}
return items;
});
// Helper to get the number of nodes in a placement preview
function getPreviewNodeCount(preview: PlacementPreview): number {
@@ -1985,31 +2069,51 @@
bind:this={instancesContainerRef}
class="max-h-72 xl:max-h-96 space-y-3 overflow-y-auto overflow-x-hidden py-px"
>
{#each Object.entries(instanceData) as [id, instance]}
{@const downloadInfo = getInstanceDownloadStatus(
id,
instance,
)}
{#each unifiedDisplayItems() as item (item.id)}
{@const id = item.id}
{@const instance = item.instance}
{@const downloadInfo = instance
? getInstanceDownloadStatus(item.instanceId ?? id, instance)
: {
statusText: "PLACING",
statusClass: "starting",
isDownloading: false,
isFailed: false,
progress: null,
perNode: [],
errorMessage: null,
}}
{@const statusText = downloadInfo.statusText}
{@const isDownloading = downloadInfo.isDownloading}
{@const isFailed = statusText === "FAILED"}
{@const isLoading =
statusText === "LOADING" ||
statusText === "WARMING UP" ||
statusText === "WAITING"}
statusText === "WAITING" ||
statusText === "PLACING"}
{@const isReady =
statusText === "READY" || statusText === "LOADED"}
{@const isRunning = statusText === "RUNNING"}
<!-- Instance Card -->
{@const instanceModelId = getInstanceModelId(instance)}
{@const instanceInfo = getInstanceInfo(instance)}
{@const instanceConnections =
getInstanceConnections(instance)}
{@const instanceModelId = item.modelId}
{@const instanceInfo = instance
? getInstanceInfo(instance)
: {
instanceType: "Unknown",
sharding: "Unknown",
nodeNames: [],
nodeIds: [],
nodeCount: 0,
}}
{@const instanceConnections = instance
? getInstanceConnections(instance)
: []}
<div
class="relative group cursor-pointer"
role="button"
tabindex="0"
onmouseenter={() => (hoveredInstanceId = id)}
onmouseenter={() =>
(hoveredInstanceId = item.instanceId ?? id)}
onmouseleave={() => (hoveredInstanceId = null)}
onclick={() => {
if (
@@ -2108,7 +2212,10 @@
>
</div>
<button
onclick={() => deleteInstance(id)}
onclick={() =>
item.isMetaInstance
? deleteMetaInstance(id)
: deleteInstance(id)}
class="text-xs px-2 py-1 font-mono tracking-wider uppercase border border-red-500/30 text-red-400 hover:bg-red-500/20 hover:text-red-400 hover:border-red-500/50 transition-all duration-200 cursor-pointer"
>
DELETE
@@ -2118,7 +2225,7 @@
<div
class="text-exo-yellow text-xs font-mono tracking-wide truncate"
>
{getInstanceModelId(instance)}
{instanceModelId}
</div>
<div class="text-white/60 text-xs font-mono">
Strategy: <span class="text-white/80"
@@ -2820,31 +2927,54 @@
<div
class="space-y-3 max-h-72 xl:max-h-96 overflow-y-auto overflow-x-hidden py-px pr-1"
>
{#each Object.entries(instanceData) as [id, instance]}
{@const downloadInfo = getInstanceDownloadStatus(
id,
instance,
)}
{#each unifiedDisplayItems() as item (item.id)}
{@const id = item.id}
{@const instance = item.instance}
{@const downloadInfo = instance
? getInstanceDownloadStatus(
item.instanceId ?? id,
instance,
)
: {
statusText: "PLACING",
statusClass: "starting",
isDownloading: false,
isFailed: false,
progress: null,
perNode: [],
errorMessage: null,
}}
{@const statusText = downloadInfo.statusText}
{@const isDownloading = downloadInfo.isDownloading}
{@const isFailed = statusText === "FAILED"}
{@const isLoading =
statusText === "LOADING" ||
statusText === "WARMING UP" ||
statusText === "WAITING"}
statusText === "WAITING" ||
statusText === "PLACING"}
{@const isReady =
statusText === "READY" || statusText === "LOADED"}
{@const isRunning = statusText === "RUNNING"}
<!-- Instance Card -->
{@const instanceModelId = getInstanceModelId(instance)}
{@const instanceInfo = getInstanceInfo(instance)}
{@const instanceConnections =
getInstanceConnections(instance)}
{@const instanceModelId = item.modelId}
{@const instanceInfo = instance
? getInstanceInfo(instance)
: {
instanceType: "Unknown",
sharding: "Unknown",
nodeNames: [],
nodeIds: [],
nodeCount: 0,
}}
{@const instanceConnections = instance
? getInstanceConnections(instance)
: []}
<div
class="relative group cursor-pointer"
role="button"
tabindex="0"
onmouseenter={() => (hoveredInstanceId = id)}
onmouseenter={() =>
(hoveredInstanceId = item.instanceId ?? id)}
onmouseleave={() => (hoveredInstanceId = null)}
onclick={() => {
if (
@@ -2943,7 +3073,10 @@
>
</div>
<button
onclick={() => deleteInstance(id)}
onclick={() =>
item.isMetaInstance
? deleteMetaInstance(id)
: deleteInstance(id)}
class="text-xs px-2 py-1 font-mono tracking-wider uppercase border border-red-500/30 text-red-400 hover:bg-red-500/20 hover:text-red-400 hover:border-red-500/50 transition-all duration-200 cursor-pointer"
>
DELETE
@@ -2953,7 +3086,7 @@
<div
class="text-exo-yellow text-xs font-mono tracking-wide truncate"
>
{getInstanceModelId(instance)}
{instanceModelId}
</div>
<div class="text-white/60 text-xs font-mono">
Strategy: <span class="text-white/80"
+46
View File
@@ -71,8 +71,11 @@ from exo.shared.types.api import (
ChatCompletionResponse,
CreateInstanceParams,
CreateInstanceResponse,
CreateMetaInstanceParams,
CreateMetaInstanceResponse,
DeleteDownloadResponse,
DeleteInstanceResponse,
DeleteMetaInstanceResponse,
ErrorInfo,
ErrorResponse,
FinishReason,
@@ -115,8 +118,10 @@ from exo.shared.types.claude_api import (
from exo.shared.types.commands import (
Command,
CreateInstance,
CreateMetaInstance,
DeleteDownload,
DeleteInstance,
DeleteMetaInstance,
DownloadCommand,
ForwarderCommand,
ForwarderDownloadCommand,
@@ -137,6 +142,7 @@ from exo.shared.types.events import (
TracesMerged,
)
from exo.shared.types.memory import Memory
from exo.shared.types.meta_instance import MetaInstance, MetaInstanceId
from exo.shared.types.openai_responses import (
ResponsesRequest,
ResponsesResponse,
@@ -275,6 +281,8 @@ class API:
self.app.get("/instance/previews")(self.get_placement_previews)
self.app.get("/instance/{instance_id}")(self.get_instance)
self.app.delete("/instance/{instance_id}")(self.delete_instance)
self.app.post("/meta_instance")(self.create_meta_instance)
self.app.delete("/meta_instance/{meta_instance_id}")(self.delete_meta_instance)
self.app.get("/models")(self.get_models)
self.app.get("/v1/models")(self.get_models)
self.app.post("/models/add")(self.add_custom_model)
@@ -521,6 +529,44 @@ class API:
instance_id=instance_id,
)
async def create_meta_instance(
self, payload: CreateMetaInstanceParams
) -> CreateMetaInstanceResponse:
meta_instance = MetaInstance(
model_id=payload.model_id,
sharding=payload.sharding,
instance_meta=payload.instance_meta,
min_nodes=payload.min_nodes,
node_ids=frozenset(payload.node_ids) if payload.node_ids else None,
)
command = CreateMetaInstance(meta_instance=meta_instance)
await self._send(command)
return CreateMetaInstanceResponse(
message="Command received.",
command_id=command.command_id,
meta_instance_id=meta_instance.meta_instance_id,
)
async def delete_meta_instance(
self, meta_instance_id: MetaInstanceId
) -> DeleteMetaInstanceResponse:
meta = self.state.meta_instances.get(meta_instance_id)
if not meta:
raise HTTPException(status_code=404, detail="MetaInstance not found")
# Delete the explicitly bound backing instance
backing_id = self.state.meta_instance_backing.get(meta_instance_id)
if backing_id:
await self._send(DeleteInstance(instance_id=backing_id))
command = DeleteMetaInstance(meta_instance_id=meta_instance_id)
await self._send(command)
return DeleteMetaInstanceResponse(
message="Command received.",
command_id=command.command_id,
meta_instance_id=meta_instance_id,
)
async def _token_chunk_stream(
self, command_id: CommandId
) -> AsyncGenerator[ErrorChunk | ToolCallChunk | TokenChunk, None]:
+38
View File
@@ -13,12 +13,14 @@ from exo.master.placement import (
place_instance,
)
from exo.master.reconcile import (
find_satisfying_instance,
find_unsatisfied_meta_instances,
instance_connections_healthy,
try_place_for_meta_instance,
)
from exo.shared.apply import apply
from exo.shared.constants import EXO_EVENT_LOG_DIR, EXO_TRACING_ENABLED
from exo.shared.models.model_cards import ModelCard
from exo.shared.types.commands import (
CreateInstance,
CreateMetaInstance,
@@ -42,6 +44,7 @@ from exo.shared.types.events import (
IndexedEvent,
InputChunkReceived,
InstanceDeleted,
MetaInstanceBound,
MetaInstanceCreated,
MetaInstanceDeleted,
NodeGatheredInfo,
@@ -294,6 +297,20 @@ class Master:
generated_events.append(
MetaInstanceCreated(meta_instance=command.meta_instance)
)
# Immediate placement attempt for responsiveness
model_card = await ModelCard.load(
command.meta_instance.model_id
)
generated_events.extend(
try_place_for_meta_instance(
command.meta_instance,
model_card,
self.state.topology,
self.state.instances,
self.state.node_memory,
self.state.node_network,
)
)
case DeleteMetaInstance():
generated_events.append(
MetaInstanceDeleted(
@@ -391,10 +408,31 @@ class Master:
self.state.meta_instances,
self.state.instances,
self.state.topology,
self.state.meta_instance_backing,
)
# Instances already bound by other MetaInstances
already_bound = frozenset(self.state.meta_instance_backing.values())
for meta_instance in unsatisfied:
# Try to bind to an existing unbound instance first
existing = find_satisfying_instance(
meta_instance,
self.state.instances,
self.state.topology,
exclude=already_bound,
)
if existing is not None:
await self._apply_and_broadcast(
MetaInstanceBound(
meta_instance_id=meta_instance.meta_instance_id,
instance_id=existing,
)
)
continue
# Otherwise, place a new instance
model_card = await ModelCard.load(meta_instance.model_id)
events = try_place_for_meta_instance(
meta_instance,
model_card,
self.state.topology,
self.state.instances,
self.state.node_memory,
+38 -12
View File
@@ -3,10 +3,11 @@ from collections.abc import Mapping, Sequence
from loguru import logger
from exo.master.placement import get_transition_events, place_instance
from exo.shared.models.model_cards import ModelCard
from exo.shared.topology import Topology
from exo.shared.types.commands import PlaceInstance
from exo.shared.types.common import NodeId
from exo.shared.types.events import Event
from exo.shared.types.events import Event, MetaInstanceBound
from exo.shared.types.meta_instance import MetaInstance, MetaInstanceId
from exo.shared.types.profiling import MemoryUsage, NodeNetworkInfo
from exo.shared.types.topology import RDMAConnection, SocketConnection
@@ -91,7 +92,7 @@ def instance_satisfies_meta_instance(
This is a pure constraint check (model, min_nodes, node_ids).
Use ``instance_connections_healthy`` separately for topology health.
"""
if instance.shard_assignments.model_id != meta_instance.model_card.model_id:
if instance.shard_assignments.model_id != meta_instance.model_id:
return False
instance_nodes = set(instance.shard_assignments.node_to_runner.keys())
@@ -108,9 +109,13 @@ def find_satisfying_instance(
meta_instance: MetaInstance,
instances: Mapping[InstanceId, Instance],
topology: Topology,
*,
exclude: frozenset[InstanceId] = frozenset(),
) -> InstanceId | None:
"""Find an existing instance that is healthy and satisfies a meta-instance's constraints."""
for instance_id, instance in instances.items():
if instance_id in exclude:
continue
if instance_connections_healthy(
instance, topology
) and instance_satisfies_meta_instance(meta_instance, instance):
@@ -122,17 +127,25 @@ def find_unsatisfied_meta_instances(
meta_instances: Mapping[MetaInstanceId, MetaInstance],
instances: Mapping[InstanceId, Instance],
topology: Topology,
meta_instance_backing: Mapping[MetaInstanceId, InstanceId],
) -> Sequence[MetaInstance]:
"""Return meta-instances that have no healthy, satisfying backing instance."""
return [
meta_instance
for meta_instance in meta_instances.values()
if find_satisfying_instance(meta_instance, instances, topology) is None
]
"""Return meta-instances whose bound backing instance is missing or unhealthy."""
unsatisfied: list[MetaInstance] = []
for meta_id, meta_instance in meta_instances.items():
bound_id = meta_instance_backing.get(meta_id)
if bound_id is not None:
bound_instance = instances.get(bound_id)
if bound_instance is not None and instance_connections_healthy(
bound_instance, topology
):
continue # bound and healthy
unsatisfied.append(meta_instance)
return unsatisfied
def try_place_for_meta_instance(
meta_instance: MetaInstance,
model_card: ModelCard,
topology: Topology,
current_instances: Mapping[InstanceId, Instance],
node_memory: Mapping[NodeId, MemoryUsage],
@@ -140,10 +153,10 @@ def try_place_for_meta_instance(
) -> Sequence[Event]:
"""Try to place an instance satisfying the meta-instance constraints.
Returns InstanceCreated events on success, empty sequence on failure.
Returns InstanceCreated + MetaInstanceBound events on success, empty sequence on failure.
"""
command = PlaceInstance(
model_card=meta_instance.model_card,
model_card=model_card,
sharding=meta_instance.sharding,
instance_meta=meta_instance.instance_meta,
min_nodes=meta_instance.min_nodes,
@@ -159,9 +172,22 @@ def try_place_for_meta_instance(
set(meta_instance.node_ids) if meta_instance.node_ids else None
),
)
return list(get_transition_events(current_instances, target_instances))
events: list[Event] = list(
get_transition_events(current_instances, target_instances)
)
# Find the newly created instance and bind it
new_instance_ids = set(target_instances.keys()) - set(current_instances.keys())
if new_instance_ids:
new_id = next(iter(new_instance_ids))
events.append(
MetaInstanceBound(
meta_instance_id=meta_instance.meta_instance_id,
instance_id=new_id,
)
)
return events
except ValueError as e:
logger.debug(
f"MetaInstance placement not possible for {meta_instance.model_card.model_id}: {e}"
f"MetaInstance placement not possible for {meta_instance.model_id}: {e}"
)
return []
+158 -8
View File
@@ -10,6 +10,7 @@ from exo.shared.topology import Topology
from exo.shared.types.common import Host, NodeId
from exo.shared.types.events import (
IndexedEvent,
InstanceCreated,
MetaInstanceCreated,
MetaInstanceDeleted,
)
@@ -84,7 +85,7 @@ def _meta_instance(
) -> MetaInstance:
return MetaInstance(
meta_instance_id=meta_instance_id or MetaInstanceId(),
model_card=_model_card(model_id),
model_id=ModelId(model_id),
min_nodes=min_nodes,
node_ids=node_ids,
)
@@ -339,7 +340,7 @@ def test_find_multiple_could_match():
def test_unsatisfied_no_meta_instances():
result = find_unsatisfied_meta_instances({}, {}, Topology())
result = find_unsatisfied_meta_instances({}, {}, Topology(), {})
assert list(result) == []
@@ -347,8 +348,12 @@ def test_unsatisfied_one_satisfied():
meta = _meta_instance()
id_a, inst_a = _instance()
topology = _topology("node-a")
# Bound via backing map
result = find_unsatisfied_meta_instances(
{meta.meta_instance_id: meta}, {id_a: inst_a}, topology
{meta.meta_instance_id: meta},
{id_a: inst_a},
topology,
{meta.meta_instance_id: id_a},
)
assert list(result) == []
@@ -358,7 +363,7 @@ def test_unsatisfied_one_not_satisfied():
id_a, inst_a = _instance("test-org/model-y")
topology = _topology("node-a")
result = find_unsatisfied_meta_instances(
{meta.meta_instance_id: meta}, {id_a: inst_a}, topology
{meta.meta_instance_id: meta}, {id_a: inst_a}, topology, {}
)
assert list(result) == [meta]
@@ -375,6 +380,7 @@ def test_unsatisfied_mix():
},
{id_a: inst_a},
topology,
{meta_satisfied.meta_instance_id: id_a},
)
assert list(result) == [meta_unsatisfied]
@@ -384,7 +390,10 @@ def test_unsatisfied_node_disconnect():
id_a, inst_a = _instance(node_ids=["node-a", "node-b"])
topology = _topology("node-a") # node-b disconnected
result = find_unsatisfied_meta_instances(
{meta.meta_instance_id: meta}, {id_a: inst_a}, topology
{meta.meta_instance_id: meta},
{id_a: inst_a},
topology,
{meta.meta_instance_id: id_a},
)
assert list(result) == [meta]
@@ -395,7 +404,10 @@ def test_unsatisfied_edge_break():
id_a, inst_a = _instance(node_ids=["node-a", "node-b"])
topology = _topology("node-a", "node-b", connect=False) # nodes present, no edges
result = find_unsatisfied_meta_instances(
{meta.meta_instance_id: meta}, {id_a: inst_a}, topology
{meta.meta_instance_id: meta},
{id_a: inst_a},
topology,
{meta.meta_instance_id: id_a},
)
assert list(result) == [meta]
@@ -406,14 +418,152 @@ def test_unsatisfied_idempotent():
meta_instances = {meta.meta_instance_id: meta}
instances: dict[InstanceId, MlxRingInstance] = {}
result_1 = list(
find_unsatisfied_meta_instances(meta_instances, instances, topology)
find_unsatisfied_meta_instances(meta_instances, instances, topology, {})
)
result_2 = list(
find_unsatisfied_meta_instances(meta_instances, instances, topology)
find_unsatisfied_meta_instances(meta_instances, instances, topology, {})
)
assert result_1 == result_2
def test_unsatisfied_exclusive_binding():
"""Two MetaInstances for the same model: one is bound, the other is unsatisfied."""
meta_a = _meta_instance("test-org/model-x")
meta_b = _meta_instance("test-org/model-x")
id_inst, inst = _instance("test-org/model-x")
topology = _topology("node-a")
# meta_a is bound to the only instance → meta_b is unsatisfied
backing = {meta_a.meta_instance_id: id_inst}
result = find_unsatisfied_meta_instances(
{
meta_a.meta_instance_id: meta_a,
meta_b.meta_instance_id: meta_b,
},
{id_inst: inst},
topology,
backing,
)
assert list(result) == [meta_b]
def test_find_satisfying_instance_exclude():
"""find_satisfying_instance should skip instances in the exclude set."""
meta = _meta_instance()
id_a, inst_a = _instance()
id_b, inst_b = _instance()
topology = _topology("node-a")
# Without exclude, first match returned
result = find_satisfying_instance(meta, {id_a: inst_a, id_b: inst_b}, topology)
assert result is not None
# Exclude the first result → should get the other one
result2 = find_satisfying_instance(
meta, {id_a: inst_a, id_b: inst_b}, topology, exclude=frozenset({result})
)
assert result2 is not None
assert result2 != result
# Exclude both → nothing found
result3 = find_satisfying_instance(
meta, {id_a: inst_a, id_b: inst_b}, topology, exclude=frozenset({id_a, id_b})
)
assert result3 is None
def test_apply_meta_instance_bound():
"""MetaInstanceBound event should populate the backing map."""
from exo.shared.types.events import MetaInstanceBound
state = State()
meta = _meta_instance()
id_a, inst_a = _instance()
# First, create meta instance and instance in state
state = apply(
state,
IndexedEvent(idx=0, event=MetaInstanceCreated(meta_instance=meta)),
)
state = apply(
state,
IndexedEvent(idx=1, event=InstanceCreated(instance=inst_a)),
)
# Bind them
state = apply(
state,
IndexedEvent(
idx=2,
event=MetaInstanceBound(
meta_instance_id=meta.meta_instance_id, instance_id=id_a
),
),
)
assert state.meta_instance_backing[meta.meta_instance_id] == id_a
def test_apply_instance_deleted_clears_backing():
"""Deleting an instance should remove it from the backing map."""
from exo.shared.types.events import InstanceDeleted, MetaInstanceBound
state = State()
meta = _meta_instance()
id_a, inst_a = _instance()
state = apply(
state,
IndexedEvent(idx=0, event=MetaInstanceCreated(meta_instance=meta)),
)
state = apply(
state,
IndexedEvent(idx=1, event=InstanceCreated(instance=inst_a)),
)
state = apply(
state,
IndexedEvent(
idx=2,
event=MetaInstanceBound(
meta_instance_id=meta.meta_instance_id, instance_id=id_a
),
),
)
assert meta.meta_instance_id in state.meta_instance_backing
# Delete the instance
state = apply(
state,
IndexedEvent(idx=3, event=InstanceDeleted(instance_id=id_a)),
)
assert meta.meta_instance_id not in state.meta_instance_backing
def test_apply_meta_instance_deleted_clears_backing():
"""Deleting a MetaInstance should remove its entry from the backing map."""
from exo.shared.types.events import MetaInstanceBound
state = State()
meta = _meta_instance()
id_a, inst_a = _instance()
state = apply(
state,
IndexedEvent(idx=0, event=MetaInstanceCreated(meta_instance=meta)),
)
state = apply(
state,
IndexedEvent(idx=1, event=InstanceCreated(instance=inst_a)),
)
state = apply(
state,
IndexedEvent(
idx=2,
event=MetaInstanceBound(
meta_instance_id=meta.meta_instance_id, instance_id=id_a
),
),
)
# Delete the meta instance
state = apply(
state,
IndexedEvent(
idx=3, event=MetaInstanceDeleted(meta_instance_id=meta.meta_instance_id)
),
)
assert meta.meta_instance_id not in state.meta_instance_backing
# --- apply handlers ---
+27 -2
View File
@@ -12,6 +12,7 @@ from exo.shared.types.events import (
InputChunkReceived,
InstanceCreated,
InstanceDeleted,
MetaInstanceBound,
MetaInstanceCreated,
MetaInstanceDeleted,
NodeDownloadProgress,
@@ -76,6 +77,8 @@ def event_apply(event: Event, state: State) -> State:
return apply_meta_instance_created(event, state)
case MetaInstanceDeleted():
return apply_meta_instance_deleted(event, state)
case MetaInstanceBound():
return apply_meta_instance_bound(event, state)
case NodeTimedOut():
return apply_node_timed_out(event, state)
case NodeDownloadProgress():
@@ -191,7 +194,14 @@ def apply_instance_deleted(event: InstanceDeleted, state: State) -> State:
new_instances: Mapping[InstanceId, Instance] = {
iid: inst for iid, inst in state.instances.items() if iid != event.instance_id
}
return state.model_copy(update={"instances": new_instances})
new_backing: Mapping[MetaInstanceId, InstanceId] = {
mid: iid
for mid, iid in state.meta_instance_backing.items()
if iid != event.instance_id
}
return state.model_copy(
update={"instances": new_instances, "meta_instance_backing": new_backing}
)
def apply_meta_instance_created(event: MetaInstanceCreated, state: State) -> State:
@@ -208,7 +218,22 @@ def apply_meta_instance_deleted(event: MetaInstanceDeleted, state: State) -> Sta
for mid, mi in state.meta_instances.items()
if mid != event.meta_instance_id
}
return state.model_copy(update={"meta_instances": new_meta})
new_backing: Mapping[MetaInstanceId, InstanceId] = {
mid: iid
for mid, iid in state.meta_instance_backing.items()
if mid != event.meta_instance_id
}
return state.model_copy(
update={"meta_instances": new_meta, "meta_instance_backing": new_backing}
)
def apply_meta_instance_bound(event: MetaInstanceBound, state: State) -> State:
new_backing: Mapping[MetaInstanceId, InstanceId] = {
**state.meta_instance_backing,
event.meta_instance_id: event.instance_id,
}
return state.model_copy(update={"meta_instance_backing": new_backing})
def apply_runner_status_updated(event: RunnerStatusUpdated, state: State) -> State:
+28
View File
@@ -9,6 +9,7 @@ from pydantic_core import PydanticUseDefault
from exo.shared.models.model_cards import ModelCard, ModelId
from exo.shared.types.common import CommandId, NodeId
from exo.shared.types.memory import Memory
from exo.shared.types.meta_instance import MetaInstanceId
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
@@ -269,6 +270,33 @@ class DeleteInstanceResponse(BaseModel):
instance_id: InstanceId
class CreateMetaInstanceParams(BaseModel):
model_id: ModelId
sharding: Sharding = Sharding.Pipeline
instance_meta: InstanceMeta = InstanceMeta.MlxRing
min_nodes: int = 1
node_ids: list[NodeId] | None = None
@field_validator("sharding", "instance_meta", mode="plain")
@classmethod
def use_default(cls, v: object):
if not v or not isinstance(v, (Sharding, InstanceMeta)):
raise PydanticUseDefault()
return v
class CreateMetaInstanceResponse(BaseModel):
message: str
command_id: CommandId
meta_instance_id: MetaInstanceId
class DeleteMetaInstanceResponse(BaseModel):
message: str
command_id: CommandId
meta_instance_id: MetaInstanceId
class AdvancedImageParams(BaseModel):
seed: Annotated[int, Field(ge=0)] | None = None
num_inference_steps: Annotated[int, Field(ge=1, le=100)] | None = None
+6
View File
@@ -77,6 +77,11 @@ class MetaInstanceDeleted(BaseEvent):
meta_instance_id: MetaInstanceId
class MetaInstanceBound(BaseEvent):
meta_instance_id: MetaInstanceId
instance_id: InstanceId
class RunnerStatusUpdated(BaseEvent):
runner_id: RunnerId
runner_status: RunnerStatus
@@ -152,6 +157,7 @@ Event = (
| InstanceDeleted
| MetaInstanceCreated
| MetaInstanceDeleted
| MetaInstanceBound
| RunnerStatusUpdated
| RunnerDeleted
| NodeTimedOut
+2 -2
View File
@@ -2,7 +2,7 @@ from typing import final
from pydantic import Field
from exo.shared.models.model_cards import ModelCard
from exo.shared.models.model_cards import ModelId
from exo.shared.types.common import Id, NodeId
from exo.shared.types.worker.instances import InstanceMeta
from exo.shared.types.worker.shards import Sharding
@@ -18,7 +18,7 @@ class MetaInstance(FrozenModel):
"""Declarative constraint: ensure an instance matching these parameters always exists."""
meta_instance_id: MetaInstanceId = Field(default_factory=MetaInstanceId)
model_card: ModelCard
model_id: ModelId
sharding: Sharding = Sharding.Pipeline
instance_meta: InstanceMeta = InstanceMeta.MlxRing
min_nodes: int = 1
+1
View File
@@ -41,6 +41,7 @@ class State(CamelCaseModel):
)
instances: Mapping[InstanceId, Instance] = {}
meta_instances: Mapping[MetaInstanceId, MetaInstance] = {}
meta_instance_backing: Mapping[MetaInstanceId, InstanceId] = {}
runners: Mapping[RunnerId, RunnerStatus] = {}
downloads: Mapping[NodeId, Sequence[DownloadProgress]] = {}
tasks: Mapping[TaskId, Task] = {}