Skip to content

Commit eccf567

Browse files
committed
Implement tests for compute manager client
1 parent 273accf commit eccf567

7 files changed

Lines changed: 225 additions & 24 deletions

File tree

alchemiscale/compute/api.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -425,7 +425,8 @@ def computemanager_register(
425425

426426
now = datetime.utcnow()
427427
cm_registration = ComputeManagerRegistration(
428-
identifier=ComputeManagerID(compute_manager_id),
428+
manager_id=compute_manager_id.manager_id,
429+
uuid=compute_manager_id.uuid,
429430
registered=now,
430431
last_status_update=now,
431432
status=ComputeManagerStatus.OK,
@@ -434,7 +435,6 @@ def computemanager_register(
434435
)
435436

436437
compute_manager_id_ = n4js.register_computemanager(cm_registration)
437-
438438
return compute_manager_id_
439439

440440

@@ -457,12 +457,13 @@ def computemanager_get_instruction(
457457
):
458458
compute_manager_id = process_compute_manager_id_string(compute_manager_id)
459459
now = datetime.utcnow()
460-
instruction = n4js.get_computemanager_instruction(
460+
instruction, payload = n4js.get_computemanager_instruction(
461461
compute_manager_id,
462462
now - timedelta(seconds=settings.ALCHEMISCALE_COMPUTE_API_FORGIVE_TIME_SECONDS),
463463
settings.ALCHEMISCALE_COMPUTE_API_MAX_FAILURES,
464464
)
465-
return instruction
465+
payload["instruction"] = str(instruction)
466+
return payload
466467

467468

468469
@router.post("/computemanager/{compute_manager_id}/update_status")

alchemiscale/compute/client.py

Lines changed: 28 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -153,8 +153,13 @@ def set_task_result(
153153
return ScopedKey.from_dict(pdr_sk)
154154

155155

156+
class AlchemiscaleComputeManagerClientError(AlchemiscaleBaseClientError): ...
157+
158+
156159
class AlchemiscaleComputeManagerClient(AlchemiscaleBaseClient):
157160

161+
_exception = AlchemiscaleComputeManagerClientError
162+
158163
def register(self, compute_manager_id: ComputeManagerID) -> ComputeManagerID:
159164
res = self._post_resource(f"/computemanager/{compute_manager_id}/register", {})
160165
return ComputeManagerID(res)
@@ -179,9 +184,29 @@ def update_status(
179184

180185
return ComputeManagerID(res)
181186

182-
def get_instruction(self, compute_manager_id: ComputeManagerID) -> ComputeManagerID:
183-
res = self._post_resource(
187+
def get_instruction(
188+
self, compute_manager_id: ComputeManagerID
189+
) -> tuple[ComputeManagerInstruction, dict]:
190+
instruction_data = self._post_resource(
184191
f"/computemanager/{compute_manager_id}/instruction",
185192
{},
186193
)
187-
return ComputeManagerInstruction(res)
194+
195+
match instruction_data:
196+
case {
197+
"instruction": "OK",
198+
"num_registered": num_registered,
199+
"compute_service_ids": ids,
200+
}:
201+
return ComputeManagerInstruction.OK, {
202+
"num_registered": num_registered,
203+
"compute_service_ids": ids,
204+
}
205+
case {"instruction": "SKIP"}:
206+
return ComputeManagerInstruction.SKIP, {}
207+
case {"instruction": "SHUTDOWN", "message": message}:
208+
return ComputeManagerInstruction.SHUTDOWN, {message: message}
209+
case _:
210+
raise self._exception(
211+
f"Received unknown instruction pattern: {instruction_data}"
212+
)

alchemiscale/compute/settings.py

Lines changed: 28 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -118,8 +118,31 @@ class Config:
118118

119119
class ComputeManagerSettings(BaseModel):
120120

121-
api_url: str
122-
name: str
123-
status_update_interval: int
124-
logfile: Path | None
125-
max_compute_services: int
121+
api_url: str | None = Field(
122+
..., description="URL of the compute API to manager services for."
123+
)
124+
identifier: str | None = Field(
125+
..., description="Identifier for the compute identity used for authentication."
126+
)
127+
key: str | None = Field(
128+
..., description="Credential for the compute identity used for authentication."
129+
)
130+
name: str = Field(
131+
...,
132+
description=(
133+
"The name to give this compute manager. This value should be distinct from all"
134+
"other compute managers."
135+
),
136+
)
137+
138+
status_update_interval: int = Field(
139+
..., description="Time in seconds to send a status update to the compute API."
140+
)
141+
logfile: Path | None = Field(..., description="File path to write logs to.")
142+
max_compute_services: int = Field(
143+
...,
144+
description="Maximum number of compute services the manager is allowed to have running at a time.",
145+
)
146+
sleep_interval: int = Field(
147+
1800, description="Time in seconds to wait for another instruction."
148+
)

alchemiscale/storage/models.py

Lines changed: 30 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -73,7 +73,7 @@ class ComputeManagerStatus(StrEnum):
7373
class ComputeManagerID(str):
7474

7575
def __init__(self, value):
76-
super().__init__(value)
76+
super().__init__()
7777

7878
parts = self.split("-")
7979

@@ -116,6 +116,35 @@ class ComputeManagerRegistration(BaseModel):
116116
detail: str
117117
saturation: float
118118

119+
def __repr__(self): # pragma: no cover
120+
return f"<ComputeManagerRegistration('{str(self)}')>"
121+
122+
def __str__(self):
123+
return "-".join([self.manager_id, self.uuid])
124+
125+
@classmethod
126+
def from_now(cls, identifier: ComputeManagerID):
127+
now = datetime.utcnow()
128+
return cls(
129+
manager_id=identifier.manager_id,
130+
uuid=identifier.uuid,
131+
last_status_update=now,
132+
status=ComputeManagerStatus.OK,
133+
detail="",
134+
saturation=0,
135+
)
136+
137+
def to_dict(self):
138+
dct = self.dict()
139+
dct["manager_id"] = str(self.manager_id)
140+
dct["uuid"] = str(self.uuid)
141+
142+
return dct
143+
144+
@classmethod
145+
def from_dict(cls, dct):
146+
return cls(**dct_)
147+
119148

120149
class TaskProvenance(BaseModel):
121150
computeserviceid: ComputeServiceID

alchemiscale/storage/statestore.py

Lines changed: 19 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1447,12 +1447,12 @@ def get_computemanager_instruction(
14471447
compute_manager_id: ComputeManagerID,
14481448
forgive_time: datetime,
14491449
max_failures: int,
1450-
) -> ComputeManagerInstruction:
1450+
) -> tuple[ComputeManagerInstruction, dict]:
14511451

14521452
manager_id, uuid = compute_manager_id.manager_id, compute_manager_id.uuid
14531453

14541454
query = """
1455-
MATCH (csm: ComputeServiceManager {manager_id: $manager_id, uuid: $uuid})
1455+
MATCH (csm: ComputeManagerRegistration {manager_id: $manager_id, uuid: $uuid})
14561456
OPTIONAL MATCH (csm)-[rel:MANAGES]->(csr: ComputeServiceRegistration)
14571457
RETURN csm, csr.identifier as csr_id
14581458
"""
@@ -1461,21 +1461,29 @@ def get_computemanager_instruction(
14611461

14621462
# no compute manager was found the given name and UUID
14631463
if len(results.records) == 0:
1464-
return ComputeManagerInstruction.SHUTDOWN
1464+
msg = "no compute manager was found the given name and UUID"
1465+
return ComputeManagerInstruction.SHUTDOWN, {"message": msg}
14651466

14661467
# TODO: very chatty, try and do this in a single query
14671468
# this would require a new state store method that requests
14681469
# failure times in bulk
1470+
csr_ids = []
14691471
for record in results.records:
1470-
csr_id = record["csr_id"]
1471-
if not self.compute_service_can_claim(
1472-
csr_id,
1473-
forgive_time,
1474-
max_failures,
1475-
):
1476-
return ComputeManagerInstruction.SKIP
1472+
if csr_id := record["csr_id"]:
1473+
if not self.compute_service_can_claim(
1474+
csr_id,
1475+
forgive_time,
1476+
max_failures,
1477+
):
1478+
return ComputeManagerInstruction.SKIP, {}
1479+
else:
1480+
break
1481+
csr_ids.append(ComputeServiceID(csr_id))
14771482

1478-
return ComputeManagerInstruction.OK
1483+
return ComputeManagerInstruction.OK, {
1484+
"num_registered": len(csr_ids),
1485+
"compute_service_ids": csr_ids,
1486+
}
14791487

14801488
def update_compute_manager_status(
14811489
self,

alchemiscale/tests/integration/compute/client/conftest.py

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -70,3 +70,27 @@ def compute_client_wrong_credential(uvicorn_server, compute_identity):
7070
identifier=compute_identity["identifier"],
7171
key="wrong credential",
7272
)
73+
74+
75+
@pytest.fixture(scope="module")
76+
def compute_manager_client(
77+
uvicorn_server,
78+
compute_identity,
79+
single_scoped_credentialed_compute,
80+
):
81+
return client.AlchemiscaleComputeManagerClient(
82+
api_url="http://127.0.0.1:8000/",
83+
# use the identifier for the single-scoped user who should have access to some things
84+
identifier=single_scoped_credentialed_compute.identifier,
85+
# all the test users are based on compute_identity who use the same password
86+
key=compute_identity["key"],
87+
)
88+
89+
90+
@pytest.fixture(scope="module")
91+
def manager_client_wrong_credential(uvicorn_server, compute_identity):
92+
return client.AlchemiscaleComputeManagerClient(
93+
api_url="http://127.0.0.1:8000/",
94+
identifier=compute_identity["identifier"],
95+
key="wrong credential",
96+
)
Lines changed: 91 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,91 @@
1+
from time import sleep
2+
3+
import pytest
4+
5+
from alchemiscale.compute import client
6+
from alchemiscale.storage.models import ComputeManagerID, ComputeManagerInstruction
7+
from alchemiscale.tests.integration.compute.utils import get_compute_settings_override
8+
9+
10+
class TestComputeManager:
11+
12+
def test_wrong_credential(
13+
self,
14+
scope_test,
15+
n4js_preloaded,
16+
manager_client_wrong_credential: client.AlchemiscaleComputeManagerClient,
17+
uvicorn_server,
18+
):
19+
with pytest.raises(client.AlchemiscaleComputeManagerClientError):
20+
manager_client_wrong_credential.get_info()
21+
22+
def test_refresh_credential(
23+
self,
24+
n4js_preloaded,
25+
compute_client: client.AlchemiscaleComputeManagerClient,
26+
uvicorn_server,
27+
):
28+
settings = get_compute_settings_override()
29+
assert compute_client._jwtoken is None
30+
compute_client._get_token()
31+
32+
token = compute_client._jwtoken
33+
assert token is not None
34+
35+
# token shouldn't change this fast
36+
compute_client.get_info()
37+
assert token == compute_client._jwtoken
38+
39+
# should change if we wait a bit
40+
sleep(settings.JWT_EXPIRE_SECONDS + 2)
41+
compute_client.get_info()
42+
assert token != compute_client._jwtoken
43+
44+
def test_api_check(
45+
self,
46+
n4js_preloaded,
47+
compute_manager_client: client.AlchemiscaleComputeManagerClient,
48+
uvicorn_server,
49+
):
50+
compute_manager_client._api_check()
51+
52+
def test_registration(
53+
self,
54+
n4js_preloaded,
55+
compute_manager_client: client.AlchemiscaleComputeManagerClient,
56+
):
57+
compute_manager_id = ComputeManagerID.from_manager_id("testmanager")
58+
returned_id = compute_manager_client.register(compute_manager_id)
59+
60+
assert compute_manager_id == returned_id
61+
62+
def test_deregistration(
63+
self,
64+
n4js_preloaded,
65+
compute_manager_client: client.AlchemiscaleComputeManagerClient,
66+
):
67+
compute_manager_id = ComputeManagerID.from_manager_id("testmanager")
68+
compute_manager_client.register(compute_manager_id)
69+
returned_id = compute_manager_client.deregister(compute_manager_id)
70+
assert compute_manager_id == returned_id
71+
72+
def test_get_instruction(
73+
self,
74+
n4js_preloaded,
75+
compute_manager_client: client.AlchemiscaleComputeManagerClient,
76+
):
77+
compute_manager_id = ComputeManagerID.from_manager_id("testmanager")
78+
compute_manager_client.register(compute_manager_id)
79+
instruction, payload = compute_manager_client.get_instruction(
80+
compute_manager_id
81+
)
82+
83+
assert instruction == ComputeManagerInstruction.OK, (instruction, payload)
84+
assert payload == {"compute_service_ids": [], "num_registered": 0}
85+
86+
def test_update_status(
87+
self,
88+
n4js_preloaded,
89+
computer_manager_client: client.AlchemiscaleComputeManagerClient,
90+
):
91+
raise NotImplementedError

0 commit comments

Comments
 (0)