diff --git a/README.md b/README.md index 4200c5f..2c0b043 100644 --- a/README.md +++ b/README.md @@ -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. diff --git a/maestro_worker_python/response.py b/maestro_worker_python/response.py index e860596..3b77c00 100644 --- a/maestro_worker_python/response.py +++ b/maestro_worker_python/response.py @@ -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.""" diff --git a/maestro_worker_python/serve.py b/maestro_worker_python/serve.py index 5a1964b..3a82aae 100644 --- a/maestro_worker_python/serve.py +++ b/maestro_worker_python/serve.py @@ -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): @@ -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") @@ -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})) @@ -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: @@ -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()} diff --git a/pyproject.toml b/pyproject.toml index c835820..8353644 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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" diff --git a/tests/test_serve.py b/tests/test_serve.py index 7e9bf9f..7ff4cd7 100644 --- a/tests/test_serve.py +++ b/tests/test_serve.py @@ -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 == [] diff --git a/uv.lock b/uv.lock index 6258f4d..96fca76 100644 --- a/uv.lock +++ b/uv.lock @@ -220,7 +220,7 @@ wheels = [ [[package]] name = "maestro-worker-python" -version = "5.1.2" +version = "5.2.0" source = { editable = "." } dependencies = [ { name = "fastapi" },