Skip to content

Commit 1106e48

Browse files
authored
Use lazy pirate pattern in tomato-job (#159)
* lazy pirate in jobs * ruff * fix test & nicer debug * higher frequency ketchup checks * one more timeout * retries in npoints handling * Fix metadata passing. * ruff & longer timeout * rewrite assert * io None * tidy up debug * lpp in manager
1 parent 32e7fd7 commit 1106e48

11 files changed

Lines changed: 202 additions & 145 deletions

File tree

src/tomato/daemon/driver.py

Lines changed: 18 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
from tomato.driverinterface_2_1 import ModelInterface as MI_2_1
2727
from tomato.drivers import driver_to_interface
2828
from tomato.models import Reply, Daemon
29+
from tomato.daemon import lpp
2930

3031
logger = logging.getLogger(__name__)
3132
ModelInterface = TypeVar("ModelInterface", MI_1_0, MI_2_0, MI_2_1)
@@ -371,25 +372,30 @@ def manager(port: int, timeout: int = 1000):
371372
logger.info("launched successfully")
372373
req = context.socket(zmq.REQ)
373374
req.connect(f"tcp://127.0.0.1:{port}")
374-
poller = zmq.Poller()
375-
poller.register(req, zmq.POLLIN)
376-
to = timeout
375+
lppargs = dict(
376+
endpoint=f"tcp://127.0.0.1:{port}",
377+
context=context,
378+
sender=sender,
379+
timeout=timeout,
380+
)
377381

378382
spawned_drivers = dict()
379383
driver_retries = defaultdict(int)
380384
component_retries = defaultdict(int)
381385

382386
while getattr(thread, "do_run"):
383-
req.send_pyobj(dict(cmd="status", sender=sender))
384-
events = dict(poller.poll(to))
385-
if req not in events:
386-
logger.warning("could not contact tomato-daemon in %d ms", to)
387-
to = to * 2
388-
continue
389-
elif to > timeout:
390-
to = timeout
387+
msg = dict(cmd="status", sender=sender)
388+
ret, req = lpp.comm(req, msg, **lppargs)
389+
if req.closed:
390+
thread.do_run = False
391+
break
392+
if ret.success:
393+
daemon: Daemon = ret.data
394+
else:
395+
logger.critical(ret.msg)
396+
thread.do_run = False
397+
break
391398

392-
daemon = req.recv_pyobj().data
393399
for driver in daemon.drvs.keys():
394400
args = [
395401
daemon.port,

src/tomato/daemon/io.py

Lines changed: 11 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -50,13 +50,13 @@ def merge_netcdfs(job: Job, snapshot=False):
5050
using the Component `role` as the group label.
5151
"""
5252
logger = logging.getLogger(f"{__name__}.merge_netcdf")
53-
logger.debug("opening datasets")
53+
logger.debug("opening datasets in '%s'", job.jobpath)
5454
datasets = []
55-
logger.debug(f"{job=}")
56-
logger.debug(f"{job.jobpath=}")
5755
for fn in Path(job.jobpath).glob("*.pkl"):
58-
with pickle.load(fn.open("rb")) as ds:
59-
datasets.append(ds)
56+
with fn.open("rb") as pkl:
57+
ds = pickle.load(pkl)
58+
if ds is not None:
59+
datasets.append(ds)
6060
logger.debug("creating a DataTree from %d groups", len(datasets))
6161
dt = xr.DataTree.from_dict({ds.attrs["role"]: ds for ds in datasets})
6262
logger.debug(f"{dt=}")
@@ -66,10 +66,9 @@ def merge_netcdfs(job: Job, snapshot=False):
6666
}
6767
dt.attrs = root_attrs
6868
outpath = job.snappath if snapshot else job.respath
69-
logger.debug("saving DataTree into '%s'", outpath)
69+
logger.debug("saving DataTree into a NetCDF file at '%s'", outpath)
7070
dt.to_netcdf(outpath)
7171
dt.close()
72-
logger.debug(f"{dt=}")
7372

7473

7574
def data_to_pickle(ds: xr.Dataset, path: Path, role: str):
@@ -81,9 +80,11 @@ def data_to_pickle(ds: xr.Dataset, path: Path, role: str):
8180
ds.attrs["role"] = role
8281
logger.debug("checking for existing pickle at '%s'", path)
8382
if path.exists():
84-
with pickle.load(path.open("rb")) as oldds:
85-
logger.debug("concatenating Dataset with existing data")
86-
ds = xr.concat([oldds, ds], dim="uts")
83+
with path.open("rb") as old:
84+
oldds = pickle.load(old)
85+
if oldds is not None:
86+
logger.debug("concatenating Dataset with existing data")
87+
ds = xr.concat([oldds, ds], dim="uts")
8788
logger.debug("dumping Dataset into pickle at '%s'", path)
8889
with path.open("wb") as out:
8990
pickle.dump(ds, out, protocol=5)

src/tomato/daemon/job.py

Lines changed: 101 additions & 98 deletions
Original file line numberDiff line numberDiff line change
@@ -25,10 +25,11 @@
2525
import zmq
2626
import psutil
2727
import sys
28+
import xarray as xr
2829

2930
from tomato.daemon.io import merge_netcdfs, data_to_pickle
30-
from tomato.daemon import jobdb
31-
from tomato.models import Pipeline, Daemon, Component, Device, Driver, Job, Reply
31+
from tomato.daemon import jobdb, lpp
32+
from tomato.models import Pipeline, Daemon, Component, Device, Driver, Job
3233
from dgbowl_schemas.tomato import to_payload
3334
from dgbowl_schemas.tomato.payload import Task
3435

@@ -56,10 +57,14 @@ def method_validate(
5657
address=cmps[cmp].address,
5758
channel=cmps[cmp].channel,
5859
)
59-
req.send_pyobj(dict(cmd="task_validate", params=params))
60-
ret = req.recv_pyobj()
61-
req.close()
60+
ret, req = lpp.comm(
61+
req,
62+
dict(cmd="task_validate", params=params),
63+
f"tcp://127.0.0.1:{drv.port}",
64+
context,
65+
)
6266
if ret.success:
67+
req.close()
6368
break
6469
else:
6570
return False
@@ -295,19 +300,15 @@ def manager(port: int, timeout: int = 500):
295300
thread = current_thread()
296301
logger.info("launched successfully")
297302
req: zmq.Socket = context.socket(zmq.REQ)
298-
req.RCVTIMEO = 1000
299303
req.connect(f"tcp://127.0.0.1:{port}")
300-
poller = zmq.Poller()
301-
poller.register(req, zmq.POLLIN)
304+
lppargs = dict(endpoint=f"tcp://127.0.0.1:{port}", context=context)
302305
while getattr(thread, "do_run"):
303306
logger.debug("tick")
304-
try:
305-
req.send_pyobj(dict(cmd="status", sender=f"{__name__}.manager"))
306-
ret: Reply = req.recv_pyobj()
307-
except zmq.ZMQError:
308-
logger.critical("could not contact tomato-daemon in 1 s", exc_info=True)
307+
msg = dict(cmd="status", sender=f"{__name__}.manager")
308+
ret, req = lpp.comm(req, msg, **lppargs)
309+
if req.closed:
309310
break
310-
if ret.success is False:
311+
elif ret.success is False:
311312
logger.critical("tomato-daemon is not running: %s", ret.msg)
312313
break
313314
daemon: Daemon = ret.data
@@ -466,81 +467,85 @@ def job_thread(
466467
thread = current_thread()
467468
sender = f"{__name__}.job_thread({thread.ident})"
468469
logger = logging.getLogger(sender)
469-
logger.debug(f"in job thread of {component.role!r}")
470-
471470
context = zmq.Context()
472471
req = context.socket(zmq.REQ)
473-
req.RCVTIMEO = 1000
474472
req.connect(f"tcp://127.0.0.1:{driver.port}")
475-
logger.info(f"job thread of {component.role!r} connected to tomato-daemon")
473+
lppargs = dict(
474+
endpoint=f"tcp://127.0.0.1:{driver.port}", context=context, sender=sender
475+
)
476+
477+
logger.info(
478+
"%s: job thread of %s attached to tomato-daemon", component.role, component.name
479+
)
476480

477481
kwargs = dict(address=component.address, channel=component.channel)
478482

479483
datapath = Path(jobpath) / f"{component.role}.pkl"
480-
logger.debug("distributing tasks:")
484+
logger.debug("%s: processing tasks on component %s", component.role, component.name)
481485
for ti, task in enumerate(tasks):
486+
taskid = f"{component.role}:{ti}"
487+
if task.task_name is not None:
488+
taskid += f":{task.task_name!r}"
482489
thread.current_task = task
483-
logger.info("processing task %s:%d", component.role, ti)
490+
logger.info("%s: processing task", taskid)
484491
while True:
485492
time.sleep(1e-1)
486493
if task.start_with_task_name is None:
487494
pass
488495
elif task.start_with_task_name in thread.started_task_names:
489496
pass
490497
else:
491-
logger.debug("waiting for task_name '%s'", task.start_with_task_name)
498+
logger.debug(
499+
"%s: waiting for task_name '%s'", taskid, task.start_with_task_name
500+
)
492501
continue
493-
logger.debug("polling component '%s' for task readiness", component.role)
494-
try:
495-
req.send_pyobj(dict(cmd="task_status", params={**kwargs}))
496-
ret = req.recv_pyobj()
497-
except zmq.ZMQError as e:
498-
logger.critical(e, exc_info=True)
499-
thread.crashed = True
500-
sys.exit(e)
502+
logger.debug(
503+
"%s: polling component %s for task readiness", taskid, component.name
504+
)
505+
ret, req = lpp.comm(
506+
req, dict(cmd="task_status", params={**kwargs}), **lppargs
507+
)
501508
if ret.success and ret.data["can_submit"]:
502509
break
503-
logger.warning("cannot submit onto component '%s', waiting", component.role)
510+
elif req.closed:
511+
thread.crashed = True
512+
sys.exit()
513+
logger.warning(
514+
"%s: cannot submit onto component %s, waiting", taskid, component.name
515+
)
504516

505-
logger.info("sending task %s:%d to component", component.role, ti)
506-
try:
507-
req.send_pyobj(dict(cmd="task_start", params={"task": task, **kwargs}))
508-
ret = req.recv_pyobj()
509-
except zmq.ZMQError as e:
510-
logger.critical(e, exc_info=True)
517+
logger.info("%s: sending task to component %s", taskid, component.name)
518+
msg = dict(cmd="task_start", params={"task": task, **kwargs})
519+
ret, req = lpp.comm(req, msg, **lppargs)
520+
if req.closed:
511521
thread.crashed = True
512-
sys.exit(e)
522+
sys.exit()
513523

514524
t0 = time.perf_counter()
515525
while True:
516526
tN = time.perf_counter()
517527
if tN - t0 > device.pollrate:
518-
logger.debug("polling task %s:%d for data", component.role, ti)
519-
try:
520-
req.send_pyobj(dict(cmd="task_data", params={**kwargs}))
521-
ret = req.recv_pyobj()
522-
except zmq.ZMQError as e:
523-
logger.critical(e, exc_info=True)
528+
logger.debug("%s: polling task for data", taskid)
529+
msg = dict(cmd="task_data", params={**kwargs})
530+
ret, req = lpp.comm(req, msg, **lppargs, timeout=5000)
531+
if req.closed:
524532
thread.crashed = True
525-
sys.exit(e)
526-
if ret.success:
527-
logger.debug("pickling received data")
528-
ds = ret.data
533+
sys.exit()
534+
elif ret.success and ret.data is not None:
535+
logger.debug("%s: pickling received data", taskid)
536+
ds: xr.Dataset = ret.data
529537
ds.attrs["tomato_Component"] = component.model_dump_json()
530538
data_to_pickle(ds, datapath, role=component.role)
531539
t0 += device.pollrate
532540

533-
logger.debug("polling task %s:%d for completion", component.role, ti)
534-
try:
535-
req.send_pyobj(dict(cmd="task_status", params={**kwargs}))
536-
ret = req.recv_pyobj()
537-
except zmq.ZMQError as e:
538-
logger.critical(e, exc_info=True)
541+
logger.debug("%s: polling task for completion", taskid)
542+
msg = dict(cmd="task_status", params={**kwargs})
543+
ret, req = lpp.comm(req, msg, **lppargs)
544+
if req.closed:
539545
thread.crashed = True
540-
sys.exit(e)
541-
542-
if ret.success and not ret.data["running"]:
543-
logger.info("task %s:%d no longer running, break", component.role, ti)
546+
sys.exit()
547+
elif ret.success and not ret.data["running"]:
548+
logger.info("%s: task no longer running, break", taskid)
544549
break
545550
elif ret.success is False:
546551
logger.critical(f"{ret=}")
@@ -550,50 +555,49 @@ def job_thread(
550555
task.stop_with_task_name is not None
551556
and task.stop_with_task_name in thread.started_task_names
552557
):
553-
logger.info("task %s:%d stop trigger met", component.role, ti)
554-
try:
555-
req.RCVTIMEO = 10000
556-
req.send_pyobj(dict(cmd="task_stop", params={**kwargs}))
557-
ret = req.recv_pyobj()
558-
req.RCVTIMEO = 1000
559-
except zmq.ZMQError as e:
560-
logger.critical(e, exc_info=True)
558+
logger.info("%s: task stop trigger met", taskid)
559+
msg = dict(cmd="task_stop", params={**kwargs})
560+
ret, req = lpp.comm(req, msg, **lppargs, timeout=5000)
561+
if req.closed:
561562
thread.crashed = True
562-
sys.exit(e)
563-
if ret.success and ret.data is not None:
564-
data_to_pickle(ret.data, datapath, role=component.role)
563+
sys.exit()
564+
elif ret.success and ret.data is not None:
565+
logger.debug("%s: pickling received data", taskid)
566+
ds: xr.Dataset = ret.data
567+
ds.attrs["tomato_Component"] = component.model_dump_json()
568+
data_to_pickle(ds, datapath, role=component.role)
565569
break
566570

567571
time.sleep(max(1e-1, (device.pollrate - (tN - t0)) / 2))
568-
logger.info("task %s:%d fetching final data", component.role, ti)
569-
try:
570-
req.send_pyobj(dict(cmd="task_data", params={**kwargs}))
571-
ret = req.recv_pyobj()
572-
except zmq.ZMQError as e:
573-
logger.critical(e, exc_info=True)
572+
logger.info("%s: task fetching final data", taskid)
573+
msg = dict(cmd="task_data", params={**kwargs})
574+
ret, req = lpp.comm(req, msg, **lppargs, timeout=5000)
575+
if req.closed:
574576
thread.crashed = True
575-
sys.exit(e)
576-
if ret.success:
577-
data_to_pickle(ret.data, datapath, role=component.role)
577+
sys.exit()
578+
elif ret.success and ret.data is not None:
579+
logger.debug("%s: pickling received data", taskid)
580+
ds: xr.Dataset = ret.data
581+
ds.attrs["tomato_Component"] = component.model_dump_json()
582+
data_to_pickle(ds, datapath, role=component.role)
578583
thread.completed_tasks.append(task)
579584
thread.current_task = None
580585

581-
logger.info("all tasks done on component '%s', resetting", component.role)
582-
try:
583-
if driver.version == "1.0":
584-
req.send_pyobj(dict(cmd="dev_reset", params={**kwargs}))
585-
else:
586-
req.send_pyobj(dict(cmd="cmp_reset", params={**kwargs}))
587-
req.RCVTIMEO = 10000
588-
ret = req.recv_pyobj()
589-
except zmq.ZMQError as e:
590-
logger.critical(e, exc_info=True)
586+
logger.info(
587+
"%s: all tasks done on component %s, resetting", component.role, component.name
588+
)
589+
if driver.version == "1.0":
590+
msg = dict(cmd="dev_reset", params={**kwargs})
591+
else:
592+
msg = dict(cmd="cmp_reset", params={**kwargs})
593+
ret, req = lpp.comm(req, msg, **lppargs, timeout=5000)
594+
if req.closed:
591595
thread.crashed = True
592-
sys.exit(e)
593-
if not ret.success:
594-
logger.warning("could not reset component '%s': %s", component.role, ret.msg)
596+
sys.exit()
597+
elif not ret.success:
598+
logger.warning("%s: could not reset component: %s", component.role, ret.msg)
595599
else:
596-
logger.info("reset of component '%s' complete", component.role)
600+
logger.info("%s: reset of component %s done", component.role, component.name)
597601
req.close()
598602

599603

@@ -613,15 +617,14 @@ def job_main_loop(
613617

614618
req = context.socket(zmq.REQ)
615619
req.connect(f"tcp://127.0.0.1:{port}")
616-
req.RCVTIMEO = 1000
620+
lppargs = dict(endpoint=f"tcp://127.0.0.1:{port}", context=context)
617621

618622
while True:
619-
req.send_pyobj(dict(cmd="status", sender=sender))
620-
try:
621-
daemon: Daemon = req.recv_pyobj().data
622-
except zmq.ZMQError as e:
623-
logger.critical(e, exc_info=True)
624-
sys.exit(e)
623+
ret, req = lpp.comm(req, dict(cmd="status", sender=sender), **lppargs)
624+
if ret.success:
625+
daemon: Daemon = ret.data
626+
else:
627+
sys.exit()
625628
if all([drv.port is not None for drv in daemon.drvs.values()]):
626629
break
627630
else:

0 commit comments

Comments
 (0)