Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 40 additions & 0 deletions docs/api.md
Original file line number Diff line number Diff line change
Expand Up @@ -163,6 +163,43 @@ JSON describing the instance to be placed.
Command acknowledgement. The instance may not be ready immediately; wait for it
to appear through `/instance/await` before sending inference requests.

### Keep a Model Running

**POST** `/deployments`

Asks the cluster to keep a model running. Whenever the cluster has no instance
of the model (its node left, it failed to start, someone deleted it), the master
places one with the same placement `/place_instance` uses. Any instance of the
model counts, however it was placed. An instance lost before it was ready makes
the next placement wait, from 10 s doubling up to 60 s; one lost after it was
ready is placed again at once. One deployment per model.

**Request body:** the same as `/place_instance`:

```json
{ "model_id": "mlx-community/Llama-3.2-1B-Instruct-4bit", "sharding": "Pipeline", "instance_meta": "MlxRing", "min_nodes": 1 }
```

**Response:** `{ "message", "command_id", "deployment_id" }`. `409` if the
model is already kept running.

### List Deployments

**GET** `/deployments`

Returns `{ "deployments": [{ "deployment": {...}, "status": ... }] }`. `status`
is `serving` (an instance of the model has every runner ready), `starting` (an
instance is on its way up), `placing` (one will be placed shortly) or
`cant_place` (no placement fits the cluster now; `deployment.placementError`
says why; it is tried again as soon as nodes or the links between them change, and every 30 s).

### Stop Keeping a Model Running

**DELETE** `/deployments/{deployment_id}`

Deletes the deployment and the instance it placed. An instance of the model
placed some other way is left running. `404` if there is no such deployment.

## 3. Models

### List Models
Expand Down Expand Up @@ -673,6 +710,9 @@ GET /instance/placement
GET /instance/{instance_id}
DELETE /instance/{instance_id}
POST /place_instance
GET /deployments
POST /deployments
DELETE /deployments/{deployment_id}

# Models
GET /models
Expand Down
68 changes: 68 additions & 0 deletions src/exo/api/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,9 @@
DeleteInstanceResponse,
DeleteTracesRequest,
DeleteTracesResponse,
DeploymentInfo,
DeploymentList,
DeploymentResponse,
ErrorInfo,
ErrorResponse,
FinishReason,
Expand Down Expand Up @@ -124,6 +127,7 @@
ResponsesResponse,
)
from exo.master.image_store import ImageStore
from exo.master.keeper import deployment_status
from exo.master.placement import place_instance as get_instance_placements
from exo.shared.apply import apply
from exo.shared.constants import (
Expand Down Expand Up @@ -156,8 +160,10 @@
AddCustomModelCard,
CancelDownload,
Command,
CreateDeployment,
CreateInstance,
DeleteCustomModelCard,
DeleteDeployment,
DeleteDownload,
DeleteInstance,
DeleteInstanceLink,
Expand All @@ -175,6 +181,7 @@
TextGeneration,
)
from exo.shared.types.common import CommandId, Id, NodeId, SystemId
from exo.shared.types.deployments import Deployment, DeploymentId
from exo.shared.types.events import (
ChunkGenerated,
Event,
Expand Down Expand Up @@ -349,6 +356,9 @@ def _setup_routes(self) -> None:
self.app.get("/instance/await", response_model=None)(self.await_instance)
self.app.get("/instance/{instance_id}")(self.get_instance)
self.app.delete("/instance/{instance_id}")(self.delete_instance)
self.app.get("/deployments")(self.list_deployments)
self.app.post("/deployments")(self.create_deployment)
self.app.delete("/deployments/{deployment_id}")(self.delete_deployment)
self.app.get("/v1/instance-links")(self.list_instance_links)
self.app.post("/v1/instance-links")(self.create_instance_link)
self.app.put("/v1/instance-links/{link_id}")(self.update_instance_link)
Expand Down Expand Up @@ -692,6 +702,64 @@ async def delete_instance(self, instance_id: InstanceId) -> DeleteInstanceRespon
instance_id=instance_id,
)

def list_deployments(self) -> DeploymentList:
return DeploymentList(
deployments=[
DeploymentInfo(
deployment=deployment,
status=deployment_status(deployment, self.state),
)
for deployment in self.state.deployments.values()
]
)

async def create_deployment(
self, payload: PlaceInstanceParams
) -> DeploymentResponse:
"""Keep a model running: the master places an instance of it whenever the cluster has none."""
existing = next(
(
deployment
for deployment in self.state.deployments.values()
if deployment.model_card.model_id == payload.model_id
),
None,
)
if existing is not None:
raise HTTPException(
status_code=409,
detail=f"{payload.model_id} is already kept running by deployment {existing.deployment_id}",
)
command = CreateDeployment(
deployment=Deployment(
deployment_id=DeploymentId(),
model_card=await ModelCard.load(payload.model_id),
sharding=payload.sharding,
instance_meta=payload.instance_meta,
min_nodes=payload.min_nodes,
)
)
await self._send(command)
return DeploymentResponse(
message="Command received.",
command_id=command.command_id,
deployment_id=command.deployment.deployment_id,
)

async def delete_deployment(
self, deployment_id: DeploymentId
) -> DeploymentResponse:
"""Stop keeping a model running, and delete the instance the keeper placed for it."""
if deployment_id not in self.state.deployments:
raise HTTPException(status_code=404, detail="Deployment not found")
command = DeleteDeployment(deployment_id=deployment_id)
await self._send(command)
return DeploymentResponse(
message="Command received.",
command_id=command.command_id,
deployment_id=deployment_id,
)

async def get_feature_flags(self) -> dict[str, bool]:
return {"disaggregation": ENABLE_DISAGGREGATION}

Expand Down
144 changes: 144 additions & 0 deletions src/exo/api/tests/test_deployments_api.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,144 @@
# pyright: reportUnusedFunction=false, reportAny=false
from typing import Any
from unittest.mock import AsyncMock, patch

from fastapi import FastAPI
from fastapi.testclient import TestClient

from exo.api.main import API
from exo.shared.models.model_cards import ModelCard, ModelId, ModelTask
from exo.shared.types.backends import Backend
from exo.shared.types.commands import CreateDeployment, DeleteDeployment
from exo.shared.types.deployments import Deployment, DeploymentId
from exo.shared.types.memory import Memory
from exo.shared.types.state import State
from exo.shared.types.worker.instances import InstanceMeta
from exo.shared.types.worker.shards import Sharding

MODEL = ModelCard(
model_id=ModelId("test/model-a"),
storage_size=Memory.from_bytes(1000),
n_layers=10,
hidden_size=30,
supports_tensor=True,
tasks=[ModelTask.TextGeneration],
backends=[Backend.MlxMetal],
)


def _make_api(state: State) -> Any:
app = FastAPI()
api = object.__new__(API)
api.app = app
api.state = state
api._send = AsyncMock() # pyright: ignore[reportPrivateUsage]
api._setup_exception_handlers() # pyright: ignore[reportPrivateUsage]
app.get("/deployments")(api.list_deployments)
app.post("/deployments")(api.create_deployment)
app.delete("/deployments/{deployment_id}")(api.delete_deployment)
return api


def _deployment() -> Deployment:
return Deployment(
deployment_id=DeploymentId(),
model_card=MODEL,
sharding=Sharding.Pipeline,
instance_meta=InstanceMeta.MlxRing,
min_nodes=1,
)


def test_create_deployment_sends_the_request_to_the_master() -> None:
api = _make_api(State())
client = TestClient(api.app)

with patch.object(ModelCard, "load", AsyncMock(return_value=MODEL)):
response = client.post(
"/deployments",
json={
"model_id": "test/model-a",
"sharding": "Tensor",
"instance_meta": "MlxJaccl",
"min_nodes": 2,
},
)

assert response.status_code == 200
api._send.assert_called_once()
command = api._send.call_args[0][0]
assert isinstance(command, CreateDeployment)
assert command.deployment.model_card == MODEL
assert command.deployment.sharding == Sharding.Tensor
assert command.deployment.instance_meta == InstanceMeta.MlxJaccl
assert command.deployment.min_nodes == 2
assert response.json()["deployment_id"] == command.deployment.deployment_id


def test_create_deployment_refuses_a_model_already_kept_running() -> None:
existing = _deployment()
api = _make_api(State(deployments={existing.deployment_id: existing}))
client = TestClient(api.app)

with patch.object(ModelCard, "load", AsyncMock(return_value=MODEL)):
response = client.post("/deployments", json={"model_id": "test/model-a"})

assert response.status_code == 409
assert existing.deployment_id in response.json()["error"]["message"]
api._send.assert_not_called()


def test_delete_deployment_sends_the_request_to_the_master() -> None:
existing = _deployment()
api = _make_api(State(deployments={existing.deployment_id: existing}))
client = TestClient(api.app)

response = client.delete(f"/deployments/{existing.deployment_id}")

assert response.status_code == 200
command = api._send.call_args[0][0]
assert isinstance(command, DeleteDeployment)
assert command.deployment_id == existing.deployment_id


def test_delete_unknown_deployment_returns_404() -> None:
api = _make_api(State())
client = TestClient(api.app)

response = client.delete(f"/deployments/{DeploymentId()}")

assert response.status_code == 404
api._send.assert_not_called()


def test_list_deployments_shows_each_with_its_status() -> None:
waiting = _deployment()
failing = _deployment().model_copy(
update={
"model_card": MODEL.model_copy(update={"model_id": ModelId("test/b")}),
"placement_error": "No cycles found with sufficient memory",
}
)
api = _make_api(
State(
deployments={
waiting.deployment_id: waiting,
failing.deployment_id: failing,
}
)
)
client = TestClient(api.app)

response = client.get("/deployments")

assert response.status_code == 200
listed = {
item["deployment"]["deploymentId"]: item
for item in response.json()["deployments"]
}
assert listed[waiting.deployment_id]["status"] == "placing"
assert listed[failing.deployment_id]["status"] == "cant_place"
assert (
listed[failing.deployment_id]["deployment"]["placementError"]
== "No cycles found with sufficient memory"
)
3 changes: 3 additions & 0 deletions src/exo/api/types/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,9 @@
from .api import DeleteInstanceResponse as DeleteInstanceResponse
from .api import DeleteTracesRequest as DeleteTracesRequest
from .api import DeleteTracesResponse as DeleteTracesResponse
from .api import DeploymentInfo as DeploymentInfo
from .api import DeploymentList as DeploymentList
from .api import DeploymentResponse as DeploymentResponse
from .api import ErrorInfo as ErrorInfo
from .api import ErrorResponse as ErrorResponse
from .api import FinishReason as FinishReason
Expand Down
16 changes: 16 additions & 0 deletions src/exo/api/types/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@

from exo.shared.models.model_cards import ModelCard, ModelId
from exo.shared.types.common import CommandId, NodeId
from exo.shared.types.deployments import Deployment, DeploymentId, DeploymentStatus
from exo.shared.types.memory import Memory
from exo.shared.types.text_generation import ReasoningDialect, ReasoningEffort
from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta
Expand Down Expand Up @@ -273,6 +274,21 @@ class PlaceInstanceParams(BaseModel):
min_nodes: int = 1


class DeploymentInfo(BaseModel):
deployment: Deployment
status: DeploymentStatus


class DeploymentList(BaseModel):
deployments: list[DeploymentInfo]


class DeploymentResponse(BaseModel):
message: str
command_id: CommandId
deployment_id: DeploymentId


class CreateInstanceParams(BaseModel):
instance: Instance

Expand Down
Loading
Loading