55"""
66
77from abc import abstractmethod
8+ from contextlib import contextmanager
89import logging
910import time
1011
1819 AlchemiscaleComputeManagerClientError ,
1920)
2021from .settings import ComputeManagerSettings , ComputeServiceSettings
22+ from ..sleep import InterruptableSleep , SleepInterrupted
2123
2224
2325class ComputeManager :
@@ -44,6 +46,7 @@ def __init__(
4446 )
4547
4648 self ._stop = False
49+ self .int_sleep = InterruptableSleep ()
4750
4851 logger = logging .getLogger ("AlchemiscaleComputeManager" )
4952 logger .setLevel (self .settings .loglevel )
@@ -75,34 +78,73 @@ def _register(self, steal=False):
7578 def _deregister (self ):
7679 self .client .deregister (self .compute_manager_id )
7780
78- def start (self , max_cycles : int | None = None , steal = False ):
79- self .logger .info (f"Starting up compute manager '{ self .settings .name } '" )
80- self ._register (steal = steal )
81- self .logger .info (f"Registered compute manager '{ self .compute_manager_id } '" )
82- self ._stop = False
81+ @contextmanager
82+ def _running (self , steal = False ):
83+ """Register this compute manager for the lifetime of the context.
84+
85+ This guarantees that if anything interrupts startup after registration,
86+ including int_sleep.clear(), the manager is still deregistered.
87+ """
88+ registered = False
89+
8390 try :
84- count = 0
85- self .logger .info ("Starting main loop" )
86- while self .cycle ():
87- count += 1
88- if max_cycles and count >= max_cycles :
89- self .logger .info ("Reached maximum number of cycles" )
90- break
91- self .logger .info (f"Sleeping for { self .settings .sleep_interval } seconds" )
92- time .sleep (self .settings .sleep_interval )
93- except Exception as e :
94- self .logger .error (f"Unknown exception raised: '{ str (e )} '" )
95- self .logger .info (f"Updating manager status to 'ERROR'" )
96- self .client .update_status (
97- self .compute_manager_id , ComputeManagerStatus .ERROR , detail = repr (e )
98- )
99- raise e
100- except KeyboardInterrupt :
101- self .logger .info ("Caught SIGINT/Keyboard interrupt." )
91+ self ._register (steal = steal )
92+ registered = True
93+
94+ self .logger .info (f"Registered compute manager '{ self .compute_manager_id } '" )
95+
96+ self ._stop = False
97+ self .int_sleep .clear ()
98+
99+ yield
100+
102101 finally :
103- self .logger .info (f"Deregistering '{ self .compute_manager_id } '" )
104- self ._deregister ()
105- self .logger .info (f"Deregistration successful" )
102+ if registered :
103+ self .logger .info (f"Deregistering '{ self .compute_manager_id } '" )
104+
105+ # kept here in case we add additional cleanup to stop later, such as other threads
106+ self .stop ()
107+
108+ self ._deregister ()
109+ self .logger .info ("Deregistration successful" )
110+
111+ def start (self , max_cycles : int | None = None , steal = False ):
112+ self .logger .info (f"Starting up compute manager '{ self .settings .name } '" )
113+
114+ with self ._running (steal = steal ):
115+ try :
116+ count = 0
117+ self .logger .info ("Starting main loop" )
118+
119+ while self .cycle ():
120+ count += 1
121+
122+ if max_cycles and count >= max_cycles :
123+ self .logger .info ("Reached maximum number of cycles" )
124+ break
125+
126+ self .logger .info (
127+ f"Sleeping for { self .settings .sleep_interval } seconds"
128+ )
129+ self .int_sleep (self .settings .sleep_interval )
130+
131+ except SleepInterrupted :
132+ self .logger .info ("Compute manager stopping." )
133+
134+ except KeyboardInterrupt :
135+ self .logger .info ("Caught SIGINT/Keyboard interrupt." )
136+
137+ except Exception as e :
138+ self .logger .error (f"Unknown exception raised: '{ str (e )} '" )
139+ self .logger .info ("Updating manager status to 'ERROR'" )
140+
141+ self .client .update_status (
142+ self .compute_manager_id ,
143+ ComputeManagerStatus .ERROR ,
144+ detail = repr (e ),
145+ )
146+
147+ raise
106148
107149 @abstractmethod
108150 def create_compute_services (self , data : dict , target : int ) -> int :
@@ -189,6 +231,7 @@ def _compute_jobs_to_create(self, num_tasks: int, num_active_services: int) -> i
189231 return jobs or 1
190232
191233 def stop (self ):
234+ self .int_sleep .interrupt ()
192235 self ._stop = True
193236
194237 def cycle (self ) -> bool :
0 commit comments