Skip to content

Commit dcde512

Browse files
committed
Test task completion and result push
1 parent c6b6cb1 commit dcde512

2 files changed

Lines changed: 97 additions & 4 deletions

File tree

alchemiscale/compute/service.py

Lines changed: 17 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -698,6 +698,9 @@ def __init__(self, settings: ComputeServiceSettings):
698698
# monitors.
699699
self._child_env = dict() # mods to child process env
700700
self._task_data = dict()
701+
self._executor_stack = ExecutorStack(self.settings.stack_size)
702+
self.tasks_claimed = 0
703+
self.tasks_finished = 0
701704
self._initialize_dag_tree()
702705
self._initialize_resource_monitors()
703706

@@ -723,6 +726,7 @@ def __init__(self, settings: ComputeServiceSettings):
723726
self.compute_service_id = ComputeServiceID.new_from_name(self.name)
724727

725728
self.int_sleep = InterruptableSleep()
729+
self._stop = False
726730
self._initialize_logger()
727731

728732
def _initialize_logger(self):
@@ -792,6 +796,7 @@ def consume_terminated_tasks(self):
792796
task_scoped_key, _ = terminating_nodes
793797
pdr = self._consume_results(task_scoped_key)
794798
self.push_result(task_scoped_key, pdr)
799+
self.tasks_finished = 1 + self.tasks_finished
795800

796801
def process_results(self):
797802
failed_tasks = set()
@@ -817,6 +822,7 @@ def process_results(self):
817822
for failed_task in failed_tasks:
818823
pdr = self._consume_results(failed_task)
819824
self.push_result(task_scoped_key, pdr)
825+
self.tasks_finished = 1 + tasks_finished
820826

821827
def stop(self):
822828
if self.has_tasks():
@@ -909,9 +915,17 @@ def cycle(self, max_tasks, max_time) -> bool:
909915
max_less_claimed = max_tasks - self.tasks_claimed
910916
n_claim = min(n_claim, max_less_claimed)
911917
tasks = self.claim_tasks(count=n_claim)
918+
919+
if tasks is None:
920+
self.logger.info("No tasks claimed. Compute API denied request.")
921+
time.sleep(self.deep_sleep_interval)
922+
return
923+
924+
self.logger.info("Claimed %d tasks", len([t for t in tasks if t is not None]))
925+
912926
for task in tasks:
913927
if task is not None:
914-
self.add_task(*task)
928+
self.add_task(task)
915929
self.tasks_claimed = 1 + self.tasks_claimed
916930

917931
for key in filter(lambda k: k[1] not in ("TERM", "ROOT"), self.next()):
@@ -925,13 +939,14 @@ def cycle(self, max_tasks, max_time) -> bool:
925939

926940
try:
927941
self._executor_stack.push(
928-
key, context, inputs, self.n_retries, env=self._child_env
942+
key, context, inputs, self.settings.n_retries, env=self._child_env
929943
)
930944
self.logger.info(f"Pushing {key[1]} to the execution stack")
931945
break
932946
except JailedKeyError:
933947
continue
934948

949+
time.sleep(self.sleep_interval)
935950
return True
936951

937952
def available_units(self) -> set[NodeKey]:

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

Lines changed: 80 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -11,8 +11,8 @@
1111
from alchemiscale.storage.statestore import Neo4jStore
1212
from alchemiscale.storage.objectstore import S3ObjectStore
1313
from alchemiscale.compute.client import AlchemiscaleComputeClientError
14-
from alchemiscale.compute.service import SynchronousComputeService
15-
from alchemiscale.compute.settings import ComputeServiceSettings
14+
from alchemiscale.compute.service import AsynchronousComputeService, SynchronousComputeService
15+
from alchemiscale.compute.settings import AsynchronousComputeServiceSettings, ComputeServiceSettings
1616

1717

1818
class TestSynchronousComputeService:
@@ -255,3 +255,81 @@ def test_missing_compute_manager(self, n4js_preloaded, service):
255255
match="Could not find ComputeManagerRegistration",
256256
):
257257
service._register()
258+
259+
class TestAsynchronousComputeService:
260+
261+
@pytest.fixture
262+
def service(self, n4js_preloaded, compute_client, tmpdir):
263+
with tmpdir.as_cwd():
264+
return AsynchronousComputeService(
265+
AsynchronousComputeServiceSettings(
266+
gpu_monitor_enabled=False,
267+
api_url=compute_client.api_url,
268+
identifier=compute_client.identifier,
269+
key=compute_client.key,
270+
name="test_compute_service",
271+
shared_basedir=Path("shared").absolute(),
272+
scratch_basedir=Path("scratch").absolute(),
273+
heartbeat_interval=1,
274+
sleep_interval=1,
275+
deep_sleep_interval=1,
276+
)
277+
)
278+
279+
def test_heartbeat(self, n4js_preloaded, service):
280+
n4js: Neo4jStore = n4js_preloaded
281+
282+
# register service; normally happens on service start, but needed
283+
# for heartbeats
284+
service._register()
285+
286+
# start up heartbeat thread
287+
heartbeat_thread = threading.Thread(target=service.heartbeat, daemon=True)
288+
heartbeat_thread.start()
289+
290+
# give time for a heartbeat
291+
time.sleep(2)
292+
293+
q = f"""
294+
match (csreg:ComputeServiceRegistration {{identifier: '{service.compute_service_id}'}})
295+
return csreg
296+
"""
297+
csreg = n4js.execute_query(q).records[0]["csreg"]
298+
299+
assert csreg["registered"] < csreg["heartbeat"]
300+
301+
# stop the service; should trigger heartbeat to stop
302+
service.stop()
303+
time.sleep(2)
304+
assert not heartbeat_thread.is_alive()
305+
306+
def test_cycle(self, n4js_preloaded, s3os_server_fresh, service):
307+
service._register()
308+
309+
q = """
310+
match (pdr:ProtocolDAGResultRef)
311+
return pdr
312+
"""
313+
314+
# preconditions
315+
protocoldagresultref = n4js_preloaded.execute_query(q)
316+
assert not protocoldagresultref.records
317+
318+
# note that non-None max_time will fail due to _start_time
319+
# never being set
320+
while service.cycle(max_tasks=1, max_time=None):
321+
pass
322+
323+
# postconditions
324+
protocoldagresultref = n4js_preloaded.execute_query(q)
325+
assert protocoldagresultref.records
326+
assert protocoldagresultref.records[0]["pdr"]["ok"] is True
327+
328+
q = """
329+
match (t:Task {status: 'complete'})
330+
return t
331+
"""
332+
333+
results = n4js_preloaded.execute_query(q)
334+
335+
assert results.records

0 commit comments

Comments
 (0)