1import logging
2from copy import deepcopy
3from socket import gethostname
4
5from celery import Celery
6from celery.events import EventReceiver
7
8from omegaml.mongoshim import mongo_url
9from omegaml.util import dict_merge
10
11logger = logging.getLogger(__name__)
12
13
14class CeleryTask:
15 """
16 A thin wrapper for a Celery.Task object
17
18 This is so that we can collect common delay arguments on the
19 .task() call
20 """
21
22 def __init__(self, task, kwargs):
23 """
24
25 Args:
26 task (Celery.Task): the celery task object
27 kwargs (dict): optional, the kwargs to pass to apply_async
28 """
29 self.task = task
30 self.kwargs = dict(kwargs)
31
32 def _apply_kwargs(self, task_kwargs, celery_kwargs):
33 # update task_kwargs from runtime's passed on kwargs
34 # update celery_kwargs to match celery routing semantics
35 task_kwargs.update(self.kwargs.get('task', {}))
36 celery_kwargs.update(self.kwargs.get('routing', {}))
37 if 'label' in celery_kwargs:
38 celery_kwargs['queue'] = celery_kwargs['label']
39 del celery_kwargs['label']
40
41 def _apply_auth(self, args, kwargs, celery_kwargs):
42 from omegaml.client.auth import AuthenticationEnv
43
44 AuthenticationEnv.active().taskauth(args, kwargs, celery_kwargs)
45
46 def apply_async(self, args=None, kwargs=None, **celery_kwargs):
47 """
48
49 Args:
50 args (tuple): the task args
51 kwargs (dict): the task kwargs, passed as task.apply_async(kwargs=kwargs)
52 celery_kwargs (dict): apply_async kwargs, passed as task.apply_async(..., **celery_kwargs)
53
54 Returns:
55 AsyncResult
56 """
57 args = args or tuple()
58 kwargs = kwargs or {}
59 self._apply_kwargs(kwargs, celery_kwargs)
60 self._apply_auth(args, kwargs, celery_kwargs)
61 return self.task.apply_async(args=args, kwargs=kwargs, **celery_kwargs)
62
63 def delay(self, *args, **kwargs):
64 """
65 submit the task with args and kwargs to pass on
66
67 This calls task.apply_async and passes on the self.kwargs.
68 """
69 return self.apply_async(args=args, kwargs=kwargs)
70
71 def signature(self, args=None, kwargs=None, immutable=False, **celery_kwargs):
72 """return the task signature with all kwargs and celery_kwargs applied"""
73 self._apply_kwargs(kwargs, celery_kwargs)
74 sig = self.task.signature(args=args, kwargs=kwargs, **celery_kwargs, immutable=immutable)
75 return sig
76
77 def run(self, *args, **kwargs):
78 return self.delay(*args, **kwargs)
79
80
[docs]
81class OmegaRuntime:
82 """
83 omegaml compute cluster gateway
84 """
85
86 def __init__(self, omega, bucket=None, defaults=None, celeryconf=None):
87 from omegaml.util import settings
88
89 self.omega = omega
90 defaults = defaults or settings()
91 self.bucket = bucket
92 self.pure_python = getattr(defaults, 'OMEGA_FORCE_PYTHON_CLIENT', False)
93 self.pure_python = self.pure_python or self._client_is_pure_python()
94 self._create_celery_app(defaults, celeryconf=celeryconf)
95 # temporary requirements, use .require() to set
96 self._require_kwargs = dict(task={}, routing={})
97 # fixed default arguments, use .require(always=True) to set
98 self._task_default_kwargs = dict(task={}, routing={})
99 # default routing label
100 self._default_label = self.celeryapp.conf.get('CELERY_DEFAULT_QUEUE')
101
102 def __repr__(self):
103 return f'OmegaRuntime({self.omega.__repr__()})'
104
105 @property
106 def auth(self):
107 return None
108
109 @property
110 def _common_kwargs(self):
111 common = deepcopy(self._task_default_kwargs)
112 common['task'].update(pure_python=self.pure_python, __bucket=self.bucket)
113 common['task'].update(self._require_kwargs['task'])
114 common['routing'].update(self._require_kwargs['routing'])
115 return common
116
117 @property
118 def _inspect(self):
119 return self.celeryapp.control.inspect()
120
121 @property
122 def is_local(self):
123 return self.celeryapp.conf['CELERY_ALWAYS_EAGER']
124
[docs]
125 def mode(self, local=None, logging=None):
126 """specify runtime modes
127
128 Args:
129 local (bool): if True, all execution will run locally, else on
130 the configured remote cluster
131 logging (bool|str|tuple): if True, will set the root logger output
132 at INFO level; a single string is the name of the logger,
133 typically a module name; a tuple (logger, level) will select
134 logger and the level. Valid levels are INFO, WARNING, ERROR,
135 CRITICAL, DEBUG
136
137 Usage::
138
139 # run all runtime tasks locally
140 om.runtime.mode(local=True)
141
142 # enable logging both in local and remote mode
143 om.runtime.mode(logging=True)
144
145 # select a specific module and level)
146 om.runtime.mode(logging=('sklearn', 'DEBUG'))
147
148 # disable logging
149 om.runtime.mode(logging=False)
150 """
151 if isinstance(local, bool):
152 self.celeryapp.conf['CELERY_ALWAYS_EAGER'] = local
153 self._task_default_kwargs['task']['__logging'] = logging
154 return self
155
156 def _create_celery_app(self, defaults, celeryconf=None):
157 # initialize celery as a runtimes
158 taskpkgs = defaults.OMEGA_CELERY_IMPORTS
159 celeryconf = dict(celeryconf or defaults.OMEGA_CELERY_CONFIG)
160 # ensure we use current value
161 celeryconf['CELERY_ALWAYS_EAGER'] = bool(defaults.OMEGA_LOCAL_RUNTIME)
162 if celeryconf['CELERY_RESULT_BACKEND'].startswith('mongodb://'):
163 celeryconf['CELERY_RESULT_BACKEND'] = mongo_url(self.omega, drop_kwargs=['uuidRepresentation'])
164 # initialize ssl configuration
165 if celeryconf.get('BROKER_USE_SSL'):
166 # celery > 5 requires ssl options to be specific
167 # https://docs.celeryq.dev/en/stable/userguide/configuration.html#std-setting-broker_use_ssl
168 # https://github.com/celery/kombu/issues/1493
169 # https://docs.python.org/dev/library/ssl.html#ssl.wrap_socket
170 # https://www.openssl.org/docs/man3.0/man3/SSL_CTX_set_default_verify_paths.html
171 # env variables:
172 # SSL_CERT_FILE, CA_CERTS_PATH
173 self._apply_broker_ssl(celeryconf)
174 self.celeryapp = Celery('omegaml')
175 self.celeryapp.config_from_object(celeryconf)
176 # needed to get it to actually load the tasks
177 # https://stackoverflow.com/a/35735471
178 self.celeryapp.autodiscover_tasks(taskpkgs, force=True)
179 self.celeryapp.finalize()
180
181 def _apply_broker_ssl(self, celeryconf):
182 # hook to apply broker ssl options
183 pass
184
185 def _client_is_pure_python(self):
186 try:
187 pass
188 except Exception as e:
189 logging.getLogger().info(e)
190 return True
191 else:
192 return False
193
194 def _sanitize_require(self, value):
195 # convert value into dict(label=value)
196 if isinstance(value, str):
197 return dict(label=value)
198 if isinstance(value, (list, tuple)):
199 return dict(*value)
200 return value
201
[docs]
202 def require(self, label=None, always=False, routing=None, task=None, logging=None, override=True, **kwargs):
203 """
204 specify requirements for the task execution
205
206 Use this to specify resource or routing requirements on the next task
207 call sent to the runtime. Any requirements will be reset after the
208 call has been submitted.
209
210 Args:
211 always (bool): if True requirements will persist across task calls. defaults to False
212 label (str): the label required by the worker to have a runtime task dispatched to it.
213 'local' is equivalent to calling self.mode(local=True).
214 task (dict): if specified applied to the task's kwargs, passed as task.apply_async(..., kwargs=task)
215 routing (dict): if specified applied to the task's routing, passed as task.apply_async(..., **routing)
216 logging (str|tuple): if specified, same as runtime.mode(logging=...)
217 override (bool): if True overrides previously set .require(), defaults to True
218 kwargs: requirements specification that the runtime understands
219
220 Usage:
221 om.runtime.require(label='gpu').model('foo').fit(...)
222
223 See Also:
224 - CeleryTask.apply_async
225 - celery.app.task.Task.apply_async, specifically the kwargs= and **options
226
227 Returns:
228 self
229 """
230 if label:
231 # avoid overriding a previous local() call by an erronous label
232 assert isinstance(label, str), "label must be valid, run om.runtime.labels() to list of active workers"
233 if label == 'local':
234 self.mode(local=True)
235 elif override:
236 self.mode(local=False)
237 # update routing, don't replace (#416)
238 routing = routing or {}
239 routing.update({'label': label or self._default_label})
240 task = task or {}
241 routing = routing or {}
242 if task or routing:
243 if not override:
244 # override not allowed, remove previously existing
245 ex_task = dict(
246 **self._task_default_kwargs['task'],
247 **self._require_kwargs['task'],
248 )
249 ex_routing = dict(
250 **self._task_default_kwargs['routing'],
251 **self._require_kwargs['routing'],
252 )
253 exists_or_none = lambda k, d: k not in d or d.get(k, False) is None
254 task = {k: v for k, v in task.items() if exists_or_none(k, ex_task)}
255 routing = {k: v for k, v in routing.items() if exists_or_none(k, ex_routing)}
256 if always:
257 self._task_default_kwargs['routing'].update(routing)
258 self._task_default_kwargs['task'].update(task)
259 else:
260 self._require_kwargs['routing'].update(routing)
261 self._require_kwargs['task'].update(task)
262 else:
263 # FIXME this does not work as expected (will only reset if both task and routing are False)
264 if not task:
265 self._require_kwargs['task'] = {}
266 if not routing:
267 self._require_kwargs['routing'] = {}
268 if logging is not None:
269 self.mode(logging=logging)
270 return self
271
[docs]
272 def model(self, modelname, require=None):
273 """
274 return a model for remote execution
275
276 Args:
277 modelname (str): the name of the object in om.models
278 require (dict): routing requirements for this job
279
280 Returns:
281 OmegaModelProxy
282 """
283 from omegaml.runtimes.proxies.modelproxy import OmegaModelProxy
284
285 self.require(**self._sanitize_require(require)) if require else None
286 return OmegaModelProxy(modelname, runtime=self)
287
[docs]
288 def job(self, jobname, require=None):
289 """
290 return a job for remote execution
291
292 Args:
293 jobname (str): the name of the object in om.jobs
294 require (dict): routing requirements for this job
295
296 Returns:
297 OmegaJobProxy
298 """
299 from omegaml.runtimes.proxies.jobproxy import OmegaJobProxy
300
301 self.require(**self._sanitize_require(require)) if require else None
302 return OmegaJobProxy(jobname, runtime=self)
303
[docs]
304 def script(self, scriptname, require=None):
305 """
306 return a script for remote execution
307
308 Args:
309 scriptname (str): the name of object in om.scripts
310 require (dict): routing requirements for this job
311
312 Returns:
313 OmegaScriptProxy
314 """
315 from omegaml.runtimes.proxies.scriptproxy import OmegaScriptProxy
316
317 self.require(**self._sanitize_require(require)) if require else None
318 return OmegaScriptProxy(scriptname, runtime=self)
319
[docs]
320 def experiment(self, experiment, provider=None, implied_run=True, recreate=False, **tracker_kwargs):
321 """set the tracking backend and experiment
322
323 Args:
324 experiment (str): the name of the experiment
325 provider (str): the name of the provider
326 tracker_kwargs (dict): additional kwargs for the tracker
327 recreate (bool): if True, recreate the experiment (i.e. drop and recreate,
328 this is useful to change the provider or other settings. All previous data will
329 be kept)
330
331 Returns:
332 OmegaTrackingProxy
333 """
334 from omegaml.runtimes.proxies.trackingproxy import OmegaTrackingProxy
335
336 # tracker implied_run means we are using the currently active run, i.e. with block will call exp.start()
337 tracker = OmegaTrackingProxy(
338 experiment, provider=provider, runtime=self, implied_run=implied_run, recreate=recreate, **tracker_kwargs
339 )
340 return tracker
341
[docs]
342 def task(self, name, **kwargs):
343 """
344 retrieve the task function from the celery instance
345
346 Args:
347 name (str): a registered celery task as ``module.tasks.task_name``
348 kwargs (dict): routing keywords to CeleryTask.apply_async
349
350 Returns:
351 CeleryTask
352 """
353 taskfn = self.celeryapp.tasks.get(name) or self.celeryapp.tasks.get(f'omega_{name}')
354 assert taskfn is not None, "cannot find task {name} in Celery runtime".format(**locals())
355 kwargs = dict_merge(self._common_kwargs, dict(routing=kwargs))
356 task = CeleryTask(taskfn, kwargs)
357 self._require_kwargs = dict(routing={}, task={})
358 return task
359
360 @property
361 def tasks(self):
362 """return registered task names
363
364 .. versionadded:: 0.18.2
365 """
366 return list(self.celeryapp.tasks.keys())
367
368 def result(self, task_id, wait=True):
369 from celery.result import AsyncResult
370
371 promise = AsyncResult(task_id, app=self.celeryapp)
372 return promise.get() if wait else promise
373
[docs]
374 def settings(self, require=None):
375 """return the runtimes's cluster settings"""
376 self.require(**require) if require else None
377 return self.task('omegaml.tasks.omega_settings').delay().get()
378
[docs]
379 def ping(self, *args, require=None, wait=True, timeout=10, **kwargs):
380 """
381 ping the runtime
382
383 Args:
384 args (tuple): task args
385 require (dict): routing requirements for this job
386 wait (bool): if True, wait for the task to return, else return
387 AsyncResult
388 timeout (int): if wait is True, the timeout in seconds, defaults to 10
389 kwargs (dict): task kwargs, as accepted by CeleryTask.apply_async
390
391 Returns:
392 * response (dict) for wait=True
393 * AsyncResult for wait=False
394 """
395 self.require(**require) if require else None
396 promise = self.task('omegaml.tasks.omega_ping').delay(*args, **kwargs)
397 return promise.get(timeout=timeout) if wait else promise
398
[docs]
399 def enable_hostqueues(self):
400 """enable a worker-specific queue on every worker host
401
402 Returns:
403 list of labels (one entry for each hostname)
404 """
405 control = self.celeryapp.control
406 inspect = control.inspect()
407 active = inspect.active()
408 queues = []
409 for worker in active.keys():
410 hostname = worker.split('@')[-1]
411 control.cancel_consumer(hostname)
412 control.add_consumer(hostname, destination=[worker])
413 queues.append(hostname)
414 return queues
415
[docs]
416 def workers(self):
417 """list of workers
418
419 Returns:
420 dict of workers => list of active tasks
421
422 See Also:
423 celery Inspect.active()
424 """
425 local_worker = {
426 gethostname(): [
427 {
428 'name': 'local',
429 'is_local': True,
430 }
431 ]
432 }
433 celery_workers = self._inspect.active() or {}
434 return dict_merge(local_worker, celery_workers)
435
[docs]
436 def queues(self):
437 """list queues
438
439 Returns:
440 dict of workers => list of queues
441
442 See Also:
443 celery Inspect.active_queues()
444 """
445 local_q = {gethostname(): [{'name': 'local', 'is_local': True}]}
446 celery_qs = self._inspect.active_queues() or {}
447 return dict_merge(local_q, celery_qs)
448
[docs]
449 def labels(self):
450 """list available labels
451
452 Returns:
453 dict of workers => list of lables
454 """
455 return {worker: [q.get('name') for q in queues] for worker, queues in self.queues().items()}
456
[docs]
457 def stats(self):
458 """worker statistics
459
460 Returns:
461 dict of workers => dict of stats
462
463 See Also:
464 celery Inspect.stats()
465 """
466 return self._inspect.stats()
467
[docs]
468 def status(self):
469 """current cluster status
470
471 This collects key information from .labels(), .stats() and the latest
472 worker heartbeat. Note that loadavg is only available if the worker has
473 recently sent a heartbeat and may not be accurate across the cluster.
474
475 Returns:
476 snapshot (dict): a snapshot of the cluster status
477 '<worker>': {
478 'loadavg': [0.0, 0.0, 0.0], # load average in % seen by the worker (1, 5, 15 min)
479 'processes': 1, # number of active worker processes
480 'concurrency': 1, # max concurrency
481 'uptime': 0, # uptime in seconds
482 'processed': Counter(task=n), # number of tasks processed
483 'queues': ['default'], # list of queues (labels) the worker is listening on
484 }
485 """
486 labels = self.labels()
487 stats = self.stats()
488 heartbeat = self.events.latest()
489 snapshot = {
490 worker: {
491 'loadavg': heartbeat.get('loadavg', []),
492 'processes': stats[worker]['pool']['processes'],
493 'concurrency': stats[worker]['pool']['max-concurrency'],
494 'uptime': stats[worker]['uptime'],
495 'processed': stats[worker]['total'],
496 'queues': labels[worker],
497 }
498 for worker in labels
499 if worker in stats
500 }
501 return snapshot
502
503 @property
504 def events(self):
505 return CeleryEventStream(self.celeryapp)
506
[docs]
507 def callback(self, script_name, always=False, **kwargs):
508 """Add a callback to a registered script
509
510 The callback will be triggered upon successful or failed
511 execution of the runtime tasks. The script syntax is::
512
513 # script.py
514 def run(om, state=None, result=None, **kwargs):
515 # state (str): 'SUCCESS'|'ERROR'
516 # result (obj): the task's serialized result
517
518 Args:
519 script_name (str): the name of the script (in om.scripts)
520 always (bool): if True always apply this callback, defaults to False
521 **kwargs: and other kwargs to pass on to the script
522
523 Returns:
524 self
525 """
526 success_sig = (
527 self
528 .script(script_name)
529 .task(as_callback=True)
530 .signature(
531 args=['SUCCESS', script_name],
532 kwargs=kwargs,
533 immutable=False,
534 )
535 )
536 error_sig = (
537 self
538 .script(script_name)
539 .task(as_callback=True)
540 .signature(
541 args=['ERROR', script_name],
542 kwargs=kwargs,
543 immutable=False,
544 )
545 )
546
547 if always:
548 self._task_default_kwargs['routing']['link'] = success_sig
549 self._task_default_kwargs['routing']['link_error'] = error_sig
550 else:
551 self._require_kwargs['routing']['link'] = success_sig
552 self._require_kwargs['routing']['link_error'] = error_sig
553 return self
554
555
556class CeleryEventStream:
557 def __init__(self, app, limit=None, timeout=5, wakeup=False):
558 self.app = app
559 self.limit = limit
560 self.timeout = timeout
561 self.wakeup = wakeup
562 self.max_size = 100
563 self.buffer = []
564
565 def handle(self, event):
566 self.buffer.append(event)
567 if len(self.buffer) > self.max_size:
568 self.buffer = self.buffer[-1 * self.max_size :]
569
570 def listen(self, handlers=None, limit=None, timeout=None):
571 # Connect to the broker using Kombu (Celery's underlying messaging system)
572 handlers = handlers or {'worker-heartbeat': self.handle}
573 limit = limit or self.limit
574 timeout = timeout or self.timeout
575 with self.app.connection() as conn:
576 # Create the EventReceiver to listen to all events
577 recv = EventReceiver(conn, handlers=handlers, app=self.app)
578 recv.capture(limit=limit, timeout=timeout, wakeup=self.wakeup)
579
580 def latest(self, timeout=None):
581 while len(self.buffer) == 0:
582 self.listen(limit=1, timeout=timeout)
583 return self.buffer[-1]
584
585
586# apply mixins
587from omegaml.runtimes.mixins.swagger import SwaggerGenerator
588from omegaml.runtimes.mixins.taskcanvas import canvas_chain, canvas_chord, canvas_group
589
590OmegaRuntime.sequence = canvas_chain
591OmegaRuntime.parallel = canvas_group
592OmegaRuntime.mapreduce = canvas_chord
593OmegaRuntime.swagger = SwaggerGenerator.build_swagger