Skip to content
Merged
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
7 changes: 7 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,13 @@ Workers return a `WorkerResponse`:
Any other field a worker returns is dropped. `internal` is the one exception, so
a worker cannot reach a caller with a field nobody agreed to.

A worker raises `ValidationError` for a bad request (400) and
`FatalWorkerError` when the process itself can no longer serve, such as after a
CUDA error that poisons the context. A fatal error still fails its own request
with 500, then `/health` and `/inference` answer 503 and the server exits, so
the orchestrator restarts the container instead of routing more jobs to it.
Raise it `from` the original error: the response carries the whole chain.

The `/health` endpoint reports the same `worker` object, supplied through
`WORKER_NAME` and `WORKER_VERSION`, and available GPU metadata in addition to
`ok`. Deployments should set `WORKER_VERSION` to the exact worker image tag.
Expand Down
4 changes: 4 additions & 0 deletions maestro_worker_python/response.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,3 +23,7 @@ class WorkerResponse(BaseModel):
class ValidationError(Exception):
def __init__(self, reason):
self.reason = reason


class FatalWorkerError(Exception):
"""This process can no longer serve; raise it from the error that broke it."""
31 changes: 28 additions & 3 deletions maestro_worker_python/serve.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@
from .kill_process import kill_child_processes, terminate_current_process
from .load_worker import load_worker
from .request_logging import register_client_safe_request_extractor
from .response import ValidationError, WorkerResponse
from .response import FatalWorkerError, ValidationError, WorkerResponse


def filter_transactions(event, hint):
Expand Down Expand Up @@ -98,14 +98,19 @@ async def lifespan(_app: FastAPI) -> AsyncIterator[None]:

error_counter = 0
lock = asyncio.Lock()
# Set once a worker raises FatalWorkerError; there is no way back to healthy.
fatal_error: str | None = None


def _error_body(exc: Exception) -> dict:
return {"error": "".join(traceback.format_exception(None, exc, exc.__traceback__))}


@app.exception_handler(500)
async def internal_exception_handler(request: Request, exc: Exception):
global error_counter
tb = "".join(traceback.format_exception(None, exc, exc.__traceback__))
try:
return JSONResponse(status_code=500, content=jsonable_encoder({"error": tb}))
return JSONResponse(status_code=500, content=jsonable_encoder(_error_body(exc)))
finally:
if error_counter > 10:
logging.error("Too many consecutive errors, shutting down worker")
Expand All @@ -115,6 +120,22 @@ async def internal_exception_handler(request: Request, exc: Exception):
error_counter += 1


@app.exception_handler(FatalWorkerError)
async def fatal_worker_error_handler(request: Request, exc: FatalWorkerError):
global fatal_error
fatal_error = str(exc)
logging.critical("Worker reported a fatal error, shutting down worker: %s", fatal_error)
try:
return JSONResponse(status_code=500, content=jsonable_encoder(_error_body(exc)))
finally:
# uvicorn shuts down gracefully, so this response is still delivered.
terminate_current_process()


def _unavailable() -> JSONResponse:
return JSONResponse(status_code=503, content={"ok": False, "error": fatal_error})


@app.exception_handler(ValidationError)
async def validation_error_handler(request: Request, exc: ValidationError):
return JSONResponse(status_code=400, content=jsonable_encoder({"error": exc.reason}))
Expand Down Expand Up @@ -153,6 +174,8 @@ def _with_worker_identity(result):
@app.post("/inference", response_model=WorkerResponse)
async def inference(request: Request):
global error_counter
if fatal_error is not None:
return _unavailable()
params = await request.json()
result = await run_in_threadpool(model.inference, input_data=params)
async with lock:
Expand All @@ -167,4 +190,6 @@ async def index(request: Request):

@app.get("/health")
async def health(request: Request):
if fatal_error is not None:
return _unavailable()
return {"ok": True, **get_health_metadata()}
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[project]
name = "maestro-worker-python"
version = "5.1.2"
version = "5.2.0"
description = "Utility to run workers on Moises/Maestro"
readme = "README.md"
requires-python = ">=3.10,<3.14"
Expand Down
57 changes: 57 additions & 0 deletions tests/test_serve.py
Original file line number Diff line number Diff line change
Expand Up @@ -284,3 +284,60 @@ def test_inference_and_health_report_the_identity_over_http(tmp_path, monkeypatc
assert inference["billable_seconds"] == 2.0
# Both surfaces read one Settings field each, so they cannot disagree.
assert health["worker"] == inference["worker"]


def _serve_failing_worker(tmp_path, monkeypatch, raised: str):
worker_path = tmp_path / "worker.py"
worker_path.write_text(
"from maestro_worker_python.response import FatalWorkerError, ValidationError\n"
"class MoisesWorker:\n"
" calls = 0\n"
" def inference(self, input_data):\n"
" MoisesWorker.calls += 1\n"
f" {raised}\n"
)
serve_module = _import_serve(monkeypatch, worker_path)
terminations: list[bool] = []
monkeypatch.setattr(serve_module, "terminate_current_process", lambda: terminations.append(True))
client = TestClient(serve_module.app, raise_server_exceptions=False)
return serve_module, client, terminations


def test_a_fatal_worker_error_fails_the_request_then_stops_the_process_serving(tmp_path, monkeypatch):
serve_module, client, terminations = _serve_failing_worker(
tmp_path,
monkeypatch,
'raise FatalWorkerError("CUDA context lost") from RuntimeError("CUDA error: an illegal memory access")',
)

failed = client.post("/inference", json={})
assert failed.status_code == 500
# The cause is what the job owner needs to debug the fault.
assert "an illegal memory access" in failed.json()["error"]
assert terminations == [True]

health = client.get("/health")
assert health.status_code == 503
assert health.json()["ok"] is False

refused = client.post("/inference", json={})
assert refused.status_code == 503
assert serve_module.model.calls == 1
assert terminations == [True]


@pytest.mark.parametrize(
("raised", "status_code"),
[
pytest.param('raise RuntimeError("shape mismatch")', 500, id="ordinary-failure"),
pytest.param('raise ValidationError("input is too big")', 400, id="validation-error"),
],
)
def test_a_recoverable_failure_keeps_the_process_serving(tmp_path, monkeypatch, raised, status_code):
serve_module, client, terminations = _serve_failing_worker(tmp_path, monkeypatch, raised)

assert client.post("/inference", json={}).status_code == status_code
assert client.get("/health").status_code == 200
assert client.post("/inference", json={}).status_code == status_code
assert serve_module.model.calls == 2
assert terminations == []
2 changes: 1 addition & 1 deletion uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading