faberon 0.2.2__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- faberon/__init__.py +3 -0
- faberon/api/__init__.py +5 -0
- faberon/api/app.py +350 -0
- faberon/api/models.py +30 -0
- faberon/cli.py +29 -0
- faberon/executor/__init__.py +5 -0
- faberon/executor/protocol.py +55 -0
- faberon/executor/slurm.py +225 -0
- faberon/ledger/__init__.py +5 -0
- faberon/ledger/ledger.py +258 -0
- faberon/schema/__init__.py +16 -0
- faberon/schema/campaign.py +42 -0
- faberon/schema/events.py +60 -0
- faberon/schema/plan.py +17 -0
- faberon/workflow/__init__.py +19 -0
- faberon/workflow/campaign.py +229 -0
- faberon/workflow/models.py +50 -0
- faberon/workflow/proposer.py +208 -0
- faberon/workflow/runtime.py +154 -0
- faberon/workflow/status.py +38 -0
- faberon/workflow/tree.py +53 -0
- faberon-0.2.2.dist-info/METADATA +12 -0
- faberon-0.2.2.dist-info/RECORD +25 -0
- faberon-0.2.2.dist-info/WHEEL +4 -0
- faberon-0.2.2.dist-info/entry_points.txt +3 -0
faberon/__init__.py
ADDED
faberon/api/__init__.py
ADDED
faberon/api/app.py
ADDED
|
@@ -0,0 +1,350 @@
|
|
|
1
|
+
"""FastAPI application factory and routes."""
|
|
2
|
+
|
|
3
|
+
import asyncio
|
|
4
|
+
import logging
|
|
5
|
+
import os
|
|
6
|
+
from collections.abc import AsyncIterator, Awaitable, Callable
|
|
7
|
+
from concurrent.futures import ThreadPoolExecutor
|
|
8
|
+
from contextlib import asynccontextmanager
|
|
9
|
+
from uuid import UUID
|
|
10
|
+
|
|
11
|
+
import dbos._recovery
|
|
12
|
+
from dbos import DBOS, DBOSConfig, SetWorkflowID
|
|
13
|
+
from fastapi import FastAPI, HTTPException, Query, Request
|
|
14
|
+
from fastapi.middleware import Middleware
|
|
15
|
+
from fastapi.responses import JSONResponse, PlainTextResponse, StreamingResponse
|
|
16
|
+
from starlette.middleware.base import BaseHTTPMiddleware
|
|
17
|
+
from starlette.requests import Request as StarletteRequest
|
|
18
|
+
from starlette.responses import Response
|
|
19
|
+
from starlette.types import ASGIApp
|
|
20
|
+
|
|
21
|
+
from .. import __version__
|
|
22
|
+
from ..executor import Executor
|
|
23
|
+
from ..executor.slurm import SlurmExecutor
|
|
24
|
+
from ..ledger import Ledger
|
|
25
|
+
from ..schema.campaign import Campaign, CampaignInfo, CampaignStatus
|
|
26
|
+
from ..schema.events import Actor, Event, EventType
|
|
27
|
+
from ..workflow import (
|
|
28
|
+
AgentProposer,
|
|
29
|
+
CampaignRunner,
|
|
30
|
+
CampaignSetup,
|
|
31
|
+
Runtime,
|
|
32
|
+
get_campaign_info,
|
|
33
|
+
)
|
|
34
|
+
from .models import CampaignCreate, CampaignCreated, CancelCampaign
|
|
35
|
+
|
|
36
|
+
_HEALTHZ_PATH = "/healthz"
|
|
37
|
+
RequestResponseEndpoint = Callable[[StarletteRequest], Awaitable[Response]]
|
|
38
|
+
|
|
39
|
+
# How long uvicorn waits before cancelling open connections at shutdown, such
|
|
40
|
+
# as an SSE stream. This would otherwise deadlock.
|
|
41
|
+
GRACEFUL_SHUTDOWN_TIMEOUT = 5
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def _rebuild_setup(campaign: Campaign) -> CampaignSetup:
|
|
45
|
+
"""Reconstruct the campaign's original inputs from its ledger row."""
|
|
46
|
+
return CampaignSetup(
|
|
47
|
+
campaign_id=campaign.campaign_id,
|
|
48
|
+
plan=campaign.plan,
|
|
49
|
+
command=campaign.command,
|
|
50
|
+
repo_path=campaign.repo_path,
|
|
51
|
+
poll_interval_seconds=campaign.poll_interval_seconds,
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
class BearerAuthMiddleware(BaseHTTPMiddleware):
|
|
56
|
+
"""Reject requests missing the expected bearer token."""
|
|
57
|
+
|
|
58
|
+
def __init__(self, app: ASGIApp, token: str) -> None:
|
|
59
|
+
super().__init__(app)
|
|
60
|
+
self._token = token
|
|
61
|
+
|
|
62
|
+
async def dispatch(
|
|
63
|
+
self, request: StarletteRequest, call_next: RequestResponseEndpoint
|
|
64
|
+
) -> Response:
|
|
65
|
+
# health checks do not need the token
|
|
66
|
+
if request.url.path == _HEALTHZ_PATH:
|
|
67
|
+
return await call_next(request)
|
|
68
|
+
auth = request.headers.get("Authorization", "")
|
|
69
|
+
if auth == f"Bearer {self._token}":
|
|
70
|
+
return await call_next(request)
|
|
71
|
+
return JSONResponse(status_code=401, content={"detail": "unauthorized"})
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def create_app(
|
|
75
|
+
executor: Executor,
|
|
76
|
+
*,
|
|
77
|
+
config_name: str = "default",
|
|
78
|
+
auth_token: str | None = None,
|
|
79
|
+
) -> FastAPI:
|
|
80
|
+
"""Build the Faberon HTTP app.
|
|
81
|
+
Requires ``FABERON_DATABASE_URL`` and ``FABERON_MODEL``.
|
|
82
|
+
|
|
83
|
+
Configures the process-global DBOS singleton.
|
|
84
|
+
"""
|
|
85
|
+
db_url = os.environ.get("FABERON_DATABASE_URL")
|
|
86
|
+
if not db_url:
|
|
87
|
+
raise RuntimeError("FABERON_DATABASE_URL is not set")
|
|
88
|
+
|
|
89
|
+
config: DBOSConfig = {
|
|
90
|
+
"name": "faberon",
|
|
91
|
+
"system_database_url": db_url,
|
|
92
|
+
"application_database_url": db_url,
|
|
93
|
+
"run_admin_server": False,
|
|
94
|
+
}
|
|
95
|
+
DBOS.destroy()
|
|
96
|
+
DBOS(config=config)
|
|
97
|
+
|
|
98
|
+
@asynccontextmanager
|
|
99
|
+
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
|
|
100
|
+
assert db_url is not None
|
|
101
|
+
original_recovery = dbos._recovery.startup_recovery_thread
|
|
102
|
+
no_recover = os.environ.get("FABERON_NO_RECOVER", "").lower() in (
|
|
103
|
+
"1",
|
|
104
|
+
"true",
|
|
105
|
+
"yes",
|
|
106
|
+
)
|
|
107
|
+
if no_recover:
|
|
108
|
+
# Serve the API without resuming pending workflows, so the
|
|
109
|
+
# operator can inspect and resume campaigns one by one.
|
|
110
|
+
dbos._recovery.startup_recovery_thread = lambda *a, **k: None # type: ignore
|
|
111
|
+
logging.getLogger(__name__).info(
|
|
112
|
+
"FABERON_NO_RECOVER set: not resuming pending workflows at launch"
|
|
113
|
+
)
|
|
114
|
+
ledger = Ledger(db_url)
|
|
115
|
+
runtime = Runtime(executor, ledger, config_name=config_name)
|
|
116
|
+
proposer = AgentProposer.from_env(ledger)
|
|
117
|
+
runner = CampaignRunner(runtime, proposer=proposer)
|
|
118
|
+
DBOS.register_instance(runtime)
|
|
119
|
+
DBOS.register_instance(runner)
|
|
120
|
+
DBOS.launch()
|
|
121
|
+
# The SSE stream polls the (sync) ledger off the event loop. Own pool:
|
|
122
|
+
# DBOS hands its executor to asyncio.to_thread as the loop default and
|
|
123
|
+
# a stream worker must not outlive DBOS.destroy() at interpreter exit.
|
|
124
|
+
sse_pool = ThreadPoolExecutor(max_workers=1, thread_name_prefix="faberon-sse-")
|
|
125
|
+
stop_sse = asyncio.Event()
|
|
126
|
+
app.state.runtime = runtime
|
|
127
|
+
app.state.runner = runner
|
|
128
|
+
app.state.ledger = ledger
|
|
129
|
+
app.state.sse_pool = sse_pool
|
|
130
|
+
app.state.stop_sse = stop_sse
|
|
131
|
+
try:
|
|
132
|
+
yield
|
|
133
|
+
finally:
|
|
134
|
+
dbos._recovery.startup_recovery_thread = original_recovery
|
|
135
|
+
stop_sse.set()
|
|
136
|
+
sse_pool.shutdown(wait=False, cancel_futures=True)
|
|
137
|
+
DBOS.destroy()
|
|
138
|
+
proposer.close()
|
|
139
|
+
ledger.close()
|
|
140
|
+
|
|
141
|
+
middleware = (
|
|
142
|
+
[Middleware(BearerAuthMiddleware, token=auth_token)]
|
|
143
|
+
if auth_token is not None
|
|
144
|
+
else []
|
|
145
|
+
)
|
|
146
|
+
app = FastAPI(
|
|
147
|
+
title="Faberon",
|
|
148
|
+
version=__version__,
|
|
149
|
+
lifespan=lifespan,
|
|
150
|
+
middleware=middleware,
|
|
151
|
+
)
|
|
152
|
+
_register_routes(app)
|
|
153
|
+
return app
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
def create_app_slurm() -> FastAPI:
|
|
157
|
+
"""Uvicorn entrypoint: Slurm executor from the environment.
|
|
158
|
+
|
|
159
|
+
Requires ``FABERON_DATABASE_URL``, ``FABERON_SLURM_ACCOUNT``,
|
|
160
|
+
``FABERON_API_TOKEN``, and ``FABERON_MODEL``.
|
|
161
|
+
Optional ``FABERON_SLURM_OUTPUT`` sets the Slurm ``--output`` path.
|
|
162
|
+
Optional ``FABERON_SLURM_GPUS`` sets the GPU count per job (default 1).
|
|
163
|
+
Optional ``FABERON_SLURM_MAX_TIME`` sets a walltime cap (minutes)..
|
|
164
|
+
"""
|
|
165
|
+
account = os.environ.get("FABERON_SLURM_ACCOUNT")
|
|
166
|
+
if not account:
|
|
167
|
+
raise RuntimeError("FABERON_SLURM_ACCOUNT is not set")
|
|
168
|
+
token = os.environ.get("FABERON_API_TOKEN")
|
|
169
|
+
if not token:
|
|
170
|
+
raise RuntimeError(
|
|
171
|
+
"FABERON_API_TOKEN is not set. On a shared login node, "
|
|
172
|
+
"localhost is reachable by other users; the API must be guarded."
|
|
173
|
+
)
|
|
174
|
+
gpus = int(os.environ.get("FABERON_SLURM_GPUS", "1"))
|
|
175
|
+
max_walltime = os.environ.get("FABERON_SLURM_MAX_TIME")
|
|
176
|
+
if max_walltime is not None:
|
|
177
|
+
max_walltime = int(max_walltime)
|
|
178
|
+
executor = SlurmExecutor(
|
|
179
|
+
account=account,
|
|
180
|
+
output=os.environ.get("FABERON_SLURM_OUTPUT"),
|
|
181
|
+
gpus=gpus,
|
|
182
|
+
max_walltime=max_walltime,
|
|
183
|
+
)
|
|
184
|
+
return create_app(executor=executor, auth_token=token)
|
|
185
|
+
|
|
186
|
+
|
|
187
|
+
def _register_routes(app: FastAPI) -> None:
|
|
188
|
+
@app.get(_HEALTHZ_PATH)
|
|
189
|
+
def healthz() -> dict[str, str]:
|
|
190
|
+
return {"status": "ok"}
|
|
191
|
+
|
|
192
|
+
@app.post("/v0/campaigns", response_model=CampaignCreated, status_code=201)
|
|
193
|
+
def create_campaign(body: CampaignCreate) -> CampaignCreated:
|
|
194
|
+
runner: CampaignRunner = app.state.runner
|
|
195
|
+
ledger: Ledger = app.state.ledger
|
|
196
|
+
campaign_id = body.campaign_id
|
|
197
|
+
workflow_id = str(campaign_id)
|
|
198
|
+
|
|
199
|
+
existing = ledger.get_campaign(campaign_id)
|
|
200
|
+
if existing is None:
|
|
201
|
+
# One active campaign per repo: a second one would interleave
|
|
202
|
+
# commits on the same checkout.
|
|
203
|
+
for other in ledger.campaigns_on_repo(body.repo_path):
|
|
204
|
+
if get_campaign_info(ledger, other).status.is_active:
|
|
205
|
+
raise HTTPException(
|
|
206
|
+
status_code=409,
|
|
207
|
+
detail=(
|
|
208
|
+
f"repo already has an active campaign: {other.campaign_id}"
|
|
209
|
+
),
|
|
210
|
+
)
|
|
211
|
+
|
|
212
|
+
setup = CampaignSetup(
|
|
213
|
+
campaign_id=campaign_id,
|
|
214
|
+
plan=body.plan,
|
|
215
|
+
command=body.command,
|
|
216
|
+
repo_path=body.repo_path,
|
|
217
|
+
poll_interval_seconds=body.poll_interval_seconds,
|
|
218
|
+
)
|
|
219
|
+
# Idempotent: DBOS dedupes on workflow id, the ledger dedupes on
|
|
220
|
+
# the campaigns row. A retry with the same campaign_id returns the
|
|
221
|
+
# existing campaign instead of creating a new one.
|
|
222
|
+
with SetWorkflowID(workflow_id):
|
|
223
|
+
handle = DBOS.start_workflow(runner.run_campaign, setup)
|
|
224
|
+
ledger.create_campaign(
|
|
225
|
+
campaign_id,
|
|
226
|
+
workflow_id,
|
|
227
|
+
body.plan,
|
|
228
|
+
body.command,
|
|
229
|
+
body.repo_path,
|
|
230
|
+
body.poll_interval_seconds,
|
|
231
|
+
)
|
|
232
|
+
return CampaignCreated(
|
|
233
|
+
campaign_id=campaign_id,
|
|
234
|
+
workflow_id=handle.workflow_id,
|
|
235
|
+
)
|
|
236
|
+
|
|
237
|
+
@app.get("/v0/campaigns")
|
|
238
|
+
def list_campaigns() -> list[CampaignInfo]:
|
|
239
|
+
ledger: Ledger = app.state.ledger
|
|
240
|
+
return [get_campaign_info(ledger, c) for c in ledger.list_campaigns()]
|
|
241
|
+
|
|
242
|
+
@app.post("/v0/campaigns/{campaign_id}/cancel", status_code=202)
|
|
243
|
+
def cancel_campaign(campaign_id: UUID, body: CancelCampaign) -> dict[str, str]:
|
|
244
|
+
ledger: Ledger = app.state.ledger
|
|
245
|
+
if ledger.get_campaign(campaign_id) is None:
|
|
246
|
+
raise HTTPException(status_code=404, detail="campaign not found")
|
|
247
|
+
events = ledger.campaign_events(campaign_id)
|
|
248
|
+
if any(e.type == EventType.CAMPAIGN_ENDED for e in events):
|
|
249
|
+
raise HTTPException(status_code=409, detail="campaign already ended")
|
|
250
|
+
ledger.append(
|
|
251
|
+
Event(
|
|
252
|
+
campaign_id=campaign_id,
|
|
253
|
+
actor=Actor.HUMAN,
|
|
254
|
+
type=EventType.CANCEL_REQUESTED,
|
|
255
|
+
justification=body.justification,
|
|
256
|
+
payload={"source": "api"},
|
|
257
|
+
)
|
|
258
|
+
)
|
|
259
|
+
DBOS.send(str(campaign_id), "cancel", "cancel")
|
|
260
|
+
return {"campaign_id": str(campaign_id), "status": "cancel requested"}
|
|
261
|
+
|
|
262
|
+
@app.post("/v0/campaigns/{campaign_id}/resume", status_code=202)
|
|
263
|
+
def resume_campaign(campaign_id: UUID) -> dict[str, str]:
|
|
264
|
+
"""Resume a PENDING campaign's workflow from its last checkpoint."""
|
|
265
|
+
runner: CampaignRunner = app.state.runner
|
|
266
|
+
ledger: Ledger = app.state.ledger
|
|
267
|
+
campaign = ledger.get_campaign(campaign_id)
|
|
268
|
+
if campaign is None:
|
|
269
|
+
raise HTTPException(status_code=404, detail="campaign not found")
|
|
270
|
+
info = get_campaign_info(ledger, campaign)
|
|
271
|
+
if info.status == CampaignStatus.ENDED:
|
|
272
|
+
raise HTTPException(
|
|
273
|
+
status_code=409, detail=f"Campaign already ended ({info.stop_reason})"
|
|
274
|
+
)
|
|
275
|
+
if info.status == CampaignStatus.DIED:
|
|
276
|
+
raise HTTPException(status_code=409, detail="Campaign workflow has died")
|
|
277
|
+
setup = _rebuild_setup(campaign)
|
|
278
|
+
with SetWorkflowID(campaign.workflow_id):
|
|
279
|
+
DBOS.start_workflow(runner.run_campaign, setup)
|
|
280
|
+
return {"campaign_id": str(campaign_id), "status": "resume requested"}
|
|
281
|
+
|
|
282
|
+
@app.get("/v0/campaigns/{campaign_id}")
|
|
283
|
+
def get_campaign(campaign_id: UUID) -> CampaignInfo:
|
|
284
|
+
ledger: Ledger = app.state.ledger
|
|
285
|
+
campaign = ledger.get_campaign(campaign_id)
|
|
286
|
+
if campaign is None:
|
|
287
|
+
raise HTTPException(status_code=404, detail="campaign not found")
|
|
288
|
+
return get_campaign_info(ledger, campaign)
|
|
289
|
+
|
|
290
|
+
@app.get("/v0/campaigns/{campaign_id}/events")
|
|
291
|
+
async def stream_events(
|
|
292
|
+
request: Request,
|
|
293
|
+
campaign_id: UUID,
|
|
294
|
+
after: int = Query(default=0, ge=0),
|
|
295
|
+
) -> StreamingResponse:
|
|
296
|
+
ledger: Ledger = request.app.state.ledger
|
|
297
|
+
if ledger.get_campaign(campaign_id) is None:
|
|
298
|
+
raise HTTPException(status_code=404, detail="campaign not found")
|
|
299
|
+
sse_pool: ThreadPoolExecutor = request.app.state.sse_pool
|
|
300
|
+
stop_sse: asyncio.Event = request.app.state.stop_sse
|
|
301
|
+
loop = asyncio.get_running_loop()
|
|
302
|
+
|
|
303
|
+
async def generate() -> AsyncIterator[str]:
|
|
304
|
+
cursor = after
|
|
305
|
+
while True:
|
|
306
|
+
if stop_sse.is_set() or await request.is_disconnected():
|
|
307
|
+
break
|
|
308
|
+
# Sync psycopg call; keep it off the event loop/DBOS executor
|
|
309
|
+
batch = await loop.run_in_executor(
|
|
310
|
+
sse_pool,
|
|
311
|
+
lambda: list(ledger.tail(campaign_id, after=cursor)),
|
|
312
|
+
)
|
|
313
|
+
if stop_sse.is_set():
|
|
314
|
+
break
|
|
315
|
+
for event in batch:
|
|
316
|
+
assert event.seq is not None
|
|
317
|
+
cursor = event.seq
|
|
318
|
+
yield f"event: ledger\ndata: {event.model_dump_json()}\n\n"
|
|
319
|
+
# Wake early on shutdown instead of sleeping through it.
|
|
320
|
+
try:
|
|
321
|
+
await asyncio.wait_for(stop_sse.wait(), timeout=0.5)
|
|
322
|
+
except TimeoutError:
|
|
323
|
+
pass
|
|
324
|
+
|
|
325
|
+
return StreamingResponse(
|
|
326
|
+
generate(),
|
|
327
|
+
media_type="text/event-stream",
|
|
328
|
+
headers={
|
|
329
|
+
"Cache-Control": "no-cache",
|
|
330
|
+
"Connection": "keep-alive",
|
|
331
|
+
"X-Accel-Buffering": "no",
|
|
332
|
+
},
|
|
333
|
+
)
|
|
334
|
+
|
|
335
|
+
@app.get("/v0/campaigns/{campaign_id}/events.jsonl")
|
|
336
|
+
def read_events_jsonl(
|
|
337
|
+
campaign_id: UUID,
|
|
338
|
+
after: int = Query(default=0, ge=0),
|
|
339
|
+
) -> PlainTextResponse:
|
|
340
|
+
"""Bounded snapshot of one campaign's events, one JSON event per line."""
|
|
341
|
+
ledger: Ledger = app.state.ledger
|
|
342
|
+
if ledger.get_campaign(campaign_id) is None:
|
|
343
|
+
raise HTTPException(status_code=404, detail="campaign not found")
|
|
344
|
+
lines = [
|
|
345
|
+
event.model_dump_json() for event in ledger.tail(campaign_id, after=after)
|
|
346
|
+
]
|
|
347
|
+
return PlainTextResponse(
|
|
348
|
+
"".join(f"{line}\n" for line in lines),
|
|
349
|
+
media_type="application/x-ndjson",
|
|
350
|
+
)
|
faberon/api/models.py
ADDED
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
"""External request and response models for the HTTP API."""
|
|
2
|
+
|
|
3
|
+
from uuid import UUID
|
|
4
|
+
|
|
5
|
+
from pydantic import BaseModel, Field
|
|
6
|
+
|
|
7
|
+
from ..schema.plan import ResearchPlan
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class CampaignCreate(BaseModel):
|
|
11
|
+
"""POST /v0/campaigns body."""
|
|
12
|
+
|
|
13
|
+
campaign_id: UUID
|
|
14
|
+
plan: ResearchPlan
|
|
15
|
+
command: list[str] = Field(min_length=1)
|
|
16
|
+
repo_path: str = Field(min_length=1)
|
|
17
|
+
poll_interval_seconds: float = Field(default=30.0, gt=0)
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class CancelCampaign(BaseModel):
|
|
21
|
+
"""POST /v0/campaigns/{id}/cancel body."""
|
|
22
|
+
|
|
23
|
+
justification: str = Field(min_length=1)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class CampaignCreated(BaseModel):
|
|
27
|
+
"""POST /v0/campaigns result: new campaign ID and DBOS workflow ID."""
|
|
28
|
+
|
|
29
|
+
campaign_id: UUID
|
|
30
|
+
workflow_id: str
|
faberon/cli.py
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
"""``faberon`` console entry point: start the control plane from anywhere."""
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
|
|
5
|
+
import uvicorn
|
|
6
|
+
|
|
7
|
+
from .api.app import GRACEFUL_SHUTDOWN_TIMEOUT, create_app_slurm
|
|
8
|
+
|
|
9
|
+
_DEFAULT_HOST = "127.0.0.1"
|
|
10
|
+
_DEFAULT_PORT = 8000
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def main() -> None:
|
|
14
|
+
"""Start the Faberon control plane with the Slurm executor.
|
|
15
|
+
|
|
16
|
+
Reads host and port from ``FABERON_HOST`` and ``FABERON_PORT`` if set.
|
|
17
|
+
Requires the same env vars as ``create_app_slurm``:
|
|
18
|
+
``FABERON_DATABASE_URL``, ``FABERON_SLURM_ACCOUNT``, ``FABERON_API_TOKEN``,
|
|
19
|
+
and ``FABERON_MODEL``.
|
|
20
|
+
"""
|
|
21
|
+
host = os.environ.get("FABERON_HOST", _DEFAULT_HOST)
|
|
22
|
+
port = int(os.environ.get("FABERON_PORT", str(_DEFAULT_PORT)))
|
|
23
|
+
uvicorn.run(
|
|
24
|
+
create_app_slurm,
|
|
25
|
+
host=host,
|
|
26
|
+
port=port,
|
|
27
|
+
factory=True,
|
|
28
|
+
timeout_graceful_shutdown=GRACEFUL_SHUTDOWN_TIMEOUT,
|
|
29
|
+
)
|
|
@@ -0,0 +1,55 @@
|
|
|
1
|
+
"""Job executor contract: submit, poll status, cancel."""
|
|
2
|
+
|
|
3
|
+
from enum import StrEnum
|
|
4
|
+
from typing import Protocol
|
|
5
|
+
|
|
6
|
+
from pydantic import BaseModel, Field, model_validator
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class JobState(StrEnum):
|
|
10
|
+
RUNNING = "running"
|
|
11
|
+
COMPLETED = "completed"
|
|
12
|
+
FAILED = "failed"
|
|
13
|
+
CANCELLED = "cancelled"
|
|
14
|
+
|
|
15
|
+
@property
|
|
16
|
+
def is_terminal(self) -> bool:
|
|
17
|
+
return self in _TERMINAL_STATES
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
_TERMINAL_STATES = frozenset((JobState.COMPLETED, JobState.FAILED, JobState.CANCELLED))
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class SubmitRequest(BaseModel):
|
|
24
|
+
"""Request to start a job."""
|
|
25
|
+
|
|
26
|
+
command: list[str] = Field(min_length=1)
|
|
27
|
+
submission_key: str = Field(min_length=1)
|
|
28
|
+
walltime: int = Field(gt=0) # in minutes
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
class JobInfo(BaseModel):
|
|
32
|
+
"""All current info of a certain job."""
|
|
33
|
+
|
|
34
|
+
job_id: str
|
|
35
|
+
state: JobState
|
|
36
|
+
exit_code: int | None = None
|
|
37
|
+
elapsed_seconds: float | None = None
|
|
38
|
+
|
|
39
|
+
@model_validator(mode="after")
|
|
40
|
+
def _terminal_implies_elapsed(self) -> JobInfo:
|
|
41
|
+
# terminal state should have elapsed_seconds set
|
|
42
|
+
if self.state.is_terminal and self.elapsed_seconds is None:
|
|
43
|
+
raise ValueError("elapsed_seconds is required on terminal jobs")
|
|
44
|
+
return self
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class Executor(Protocol):
|
|
48
|
+
def submit(self, request: SubmitRequest) -> str:
|
|
49
|
+
"""Start a job. Returns a job id. Same submission_key returns the same id."""
|
|
50
|
+
|
|
51
|
+
def status(self, job_id: str) -> JobInfo:
|
|
52
|
+
"""Return the current status of a job."""
|
|
53
|
+
|
|
54
|
+
def cancel(self, job_id: str) -> None:
|
|
55
|
+
"""Cancel a job if it is still running."""
|