|
9 | 9 | from datetime import timedelta |
10 | 10 | import random |
11 | 11 |
|
12 | | -from fastapi import FastAPI, APIRouter, Body, Depends, Request |
| 12 | +from fastapi import FastAPI, APIRouter, Body, Depends, HTTPException, Request |
| 13 | +from fastapi import status as http_status |
13 | 14 | from fastapi.middleware.gzip import GZipMiddleware |
14 | 15 | from gufe.tokenization import JSON_HANDLER |
15 | 16 | from gufe.protocols import ProtocolDAGResult |
|
40 | 41 | from ..storage.models import ( |
41 | 42 | ProtocolDAGResultRef, |
42 | 43 | ComputeServiceID, |
| 44 | + ComputeManagerID, |
| 45 | + ComputeManagerRegistration, |
43 | 46 | ComputeServiceRegistration, |
| 47 | + ComputeManagerStatus, |
44 | 48 | ) |
45 | 49 | from ..models import Scope, ScopedKey |
46 | 50 | from ..security.models import ( |
@@ -102,17 +106,31 @@ def list_scopes( |
102 | 106 | @router.post("/computeservice/{compute_service_id}/register") |
103 | 107 | def register_computeservice( |
104 | 108 | compute_service_id, |
| 109 | + *, |
| 110 | + compute_manager_id: str | None = Body(None, embed=True), |
105 | 111 | n4js: Neo4jStore = Depends(get_n4js_depends), |
106 | 112 | ): |
107 | 113 | now = datetime.datetime.now(tz=datetime.UTC) |
| 114 | + if compute_manager_id: |
| 115 | + manager_name = process_compute_manager_id_string(compute_manager_id).name |
| 116 | + else: |
| 117 | + manager_name = None |
| 118 | + |
108 | 119 | csreg = ComputeServiceRegistration( |
109 | 120 | identifier=ComputeServiceID(compute_service_id), |
110 | 121 | registered=now, |
111 | 122 | heartbeat=now, |
112 | 123 | failure_times=[], |
| 124 | + manager_name=manager_name, |
113 | 125 | ) |
114 | 126 |
|
115 | | - compute_service_id_ = n4js.register_computeservice(csreg) |
| 127 | + try: |
| 128 | + compute_service_id_ = n4js.register_computeservice(csreg) |
| 129 | + except ValueError as e: |
| 130 | + raise HTTPException( |
| 131 | + status_code=http_status.HTTP_422_UNPROCESSABLE_ENTITY, |
| 132 | + detail=str(e), |
| 133 | + ) |
116 | 134 |
|
117 | 135 | return compute_service_id_ |
118 | 136 |
|
@@ -392,6 +410,136 @@ async def set_task_result( |
392 | 410 | return result_sk |
393 | 411 |
|
394 | 412 |
|
| 413 | +def process_compute_manager_id_string( |
| 414 | + compute_manager_id_string: str, |
| 415 | +) -> ComputeManagerID: |
| 416 | + """Try creating a ComputeManagerID from a string representation. Raise HTTPException.""" |
| 417 | + try: |
| 418 | + compute_manager_id = ComputeManagerID(compute_manager_id_string) |
| 419 | + except Exception as e: |
| 420 | + raise HTTPException( |
| 421 | + status_code=http_status.HTTP_422_UNPROCESSABLE_ENTITY, |
| 422 | + detail=str(e), |
| 423 | + ) |
| 424 | + |
| 425 | + return compute_manager_id |
| 426 | + |
| 427 | + |
| 428 | +@router.post("/computemanager/{compute_manager_id}/register") |
| 429 | +def register_computemanager( |
| 430 | + compute_manager_id, |
| 431 | + n4js: Neo4jStore = Depends(get_n4js_depends), |
| 432 | +): |
| 433 | + |
| 434 | + compute_manager_id = process_compute_manager_id_string(compute_manager_id) |
| 435 | + |
| 436 | + now = datetime.datetime.now(tz=datetime.UTC) |
| 437 | + cm_registration = ComputeManagerRegistration( |
| 438 | + name=compute_manager_id.name, |
| 439 | + uuid=compute_manager_id.uuid, |
| 440 | + registered=now, |
| 441 | + last_status_update=now, |
| 442 | + status=ComputeManagerStatus.OK, |
| 443 | + detail="", |
| 444 | + saturation=0, |
| 445 | + ) |
| 446 | + |
| 447 | + try: |
| 448 | + compute_manager_id_ = n4js.register_computemanager(cm_registration) |
| 449 | + except ValueError as e: |
| 450 | + raise HTTPException( |
| 451 | + status_code=http_status.HTTP_422_UNPROCESSABLE_ENTITY, |
| 452 | + detail=str(e), |
| 453 | + ) |
| 454 | + |
| 455 | + return compute_manager_id_ |
| 456 | + |
| 457 | + |
| 458 | +@router.post("/computemanager/{compute_manager_id}/deregister") |
| 459 | +def deregister_computemanager( |
| 460 | + compute_manager_id, |
| 461 | + n4js: Neo4jStore = Depends(get_n4js_depends), |
| 462 | +): |
| 463 | + compute_manager_id = process_compute_manager_id_string(compute_manager_id) |
| 464 | + n4js.deregister_computemanager(compute_manager_id) |
| 465 | + return compute_manager_id |
| 466 | + |
| 467 | + |
| 468 | +@router.post("/computemanager/{compute_manager_id}/instruction") |
| 469 | +def get_instruction_computemanager( |
| 470 | + compute_manager_id, |
| 471 | + *, |
| 472 | + scopes: list[Scope] = Body([], embed=True), |
| 473 | + n4js: Neo4jStore = Depends(get_n4js_depends), |
| 474 | + settings: ComputeAPISettings = Depends(get_base_api_settings), |
| 475 | + token: TokenData = Depends(get_token_data_depends), |
| 476 | +): |
| 477 | + scopes = scopes or [Scope()] |
| 478 | + scopes_reduced = minimize_scope_space(scopes) |
| 479 | + query_scopes = [] |
| 480 | + for scope in scopes_reduced: |
| 481 | + query_scopes.extend(validate_scopes_query(scope, token)) |
| 482 | + |
| 483 | + compute_manager_id = process_compute_manager_id_string(compute_manager_id) |
| 484 | + now = datetime.datetime.now(tz=datetime.UTC) |
| 485 | + instruction, payload = n4js.get_computemanager_instruction( |
| 486 | + compute_manager_id, |
| 487 | + now - timedelta(seconds=settings.ALCHEMISCALE_COMPUTE_API_FORGIVE_TIME_SECONDS), |
| 488 | + settings.ALCHEMISCALE_COMPUTE_API_MAX_FAILURES, |
| 489 | + query_scopes, |
| 490 | + ) |
| 491 | + payload["instruction"] = str(instruction) |
| 492 | + return payload |
| 493 | + |
| 494 | + |
| 495 | +@router.post("/computemanager/{compute_manager_id}/status") |
| 496 | +def update_status_computemanager( |
| 497 | + compute_manager_id, |
| 498 | + *, |
| 499 | + status: str = Body(), |
| 500 | + detail: str | None = Body(None), |
| 501 | + saturation: float | None = Body(None), |
| 502 | + n4js: Neo4jStore = Depends(get_n4js_depends), |
| 503 | + settings: ComputeAPISettings = Depends(get_base_api_settings), |
| 504 | +): |
| 505 | + expire_seconds = settings.ALCHEMISCALE_COMPUTE_API_MANAGER_EXPIRE_SECONDS |
| 506 | + expire_seconds_errored = ( |
| 507 | + settings.ALCHEMISCALE_COMPUTE_API_MANAGER_EXPIRE_SECONDS_ERROR |
| 508 | + ) |
| 509 | + compute_manager_id = process_compute_manager_id_string(compute_manager_id) |
| 510 | + try: |
| 511 | + n4js.update_compute_manager_status( |
| 512 | + compute_manager_id, status, detail, saturation |
| 513 | + ) |
| 514 | + now = datetime.datetime.now(tz=datetime.UTC) |
| 515 | + n4js.expire_computemanager_registrations( |
| 516 | + now - timedelta(seconds=expire_seconds), |
| 517 | + now - timedelta(seconds=expire_seconds_errored), |
| 518 | + ) |
| 519 | + except ValueError as e: |
| 520 | + raise HTTPException( |
| 521 | + status_code=http_status.HTTP_400_BAD_REQUEST, |
| 522 | + detail=str(e), |
| 523 | + ) |
| 524 | + |
| 525 | + return compute_manager_id |
| 526 | + |
| 527 | + |
| 528 | +@router.post("/computemanager/{compute_manager_name}/clear_error") |
| 529 | +def clear_error_computemanager( |
| 530 | + compute_manager_name: str, |
| 531 | + n4js: Neo4jStore = Depends(get_n4js_depends), |
| 532 | +): |
| 533 | + if not compute_manager_name.isalnum(): |
| 534 | + raise ValueError("Provided manager name is not alphanumeric") |
| 535 | + |
| 536 | + with n4js.transaction() as tx: |
| 537 | + compute_manager_id = n4js.get_compute_manager_id( |
| 538 | + name=compute_manager_name, tx=tx |
| 539 | + ) |
| 540 | + n4js.clear_errored_computemanager(compute_manager_id, tx=tx) |
| 541 | + |
| 542 | + |
395 | 543 | ### add router |
396 | 544 |
|
397 | 545 | app.include_router(router) |
0 commit comments