diff --git a/src/cuphoton/core/dragon.py b/src/cuphoton/core/dragon.py index e29790fe..eea6de0d 100644 --- a/src/cuphoton/core/dragon.py +++ b/src/cuphoton/core/dragon.py @@ -403,6 +403,10 @@ def run_dragon_work_items( results_queue, ), policy=policy, + # Native Dragon workers use their own transport. Keep Python + # multiprocessing unpatched so local spawn children receive + # ordinary queues and can complete their bootstrap. + env={"DRAGON_PATCH_MP": ""}, ) if len(template.argdata) > _TEMPLATE_BUDGET_BYTES: raise ValueError( diff --git a/tests/core/test_dragon.py b/tests/core/test_dragon.py index 21473467..34cecebf 100644 --- a/tests/core/test_dragon.py +++ b/tests/core/test_dragon.py @@ -6,8 +6,11 @@ from __future__ import annotations +import importlib.util import json +import os import queue +import subprocess import sys import threading import time @@ -156,8 +159,9 @@ def __init__(self, **kwargs): self.__dict__.update(kwargs) class Template: - def __init__(self, target, args, policy): + def __init__(self, target, args, policy, env=None): self.target, self.args, self.policy = target, args, policy + self.env = env self.argdata = repr(args).encode() class Group: @@ -337,6 +341,40 @@ def finalizer(round_dir, records): assert len(run_ids) == 3 +def test_native_workers_keep_spawn_children_on_standard_multiprocessing( + monkeypatch, tmp_path +): + if importlib.util.find_spec("dragon") is None: + pytest.skip("Dragon is not installed") + monkeypatch.setenv("DRAGON_PATCH_MP", "True") + state = _install_runtime(monkeypatch) + result = _run(tmp_path) + assert result.status == "success", result.summary + assert all(channel.closed for channel in state.queues) + assert os.environ["DRAGON_PATCH_MP"] == "True" + + # Import real Dragon before creating the worker's local spawn queue. + # Without the overlay, Dragon replaces that queue with its native type. + template = state.groups[0].templates[0] + completed = subprocess.run( + [ + sys.executable, + "-c", + "import dragon; import multiprocessing as mp; " + "ctx = mp.get_context('spawn'); channel = ctx.Queue(); " + "assert type(channel).__module__ == 'multiprocessing.queues'; " + "child = ctx.Process(target=channel.put, args=('received',)); " + "child.start(); assert channel.get(timeout=5) == 'received'; " + "child.join(5); assert child.exitcode == 0; channel.close()", + ], + env={**os.environ, **(template.env or {})}, + capture_output=True, + text=True, + timeout=15, + ) + assert completed.returncode == 0, completed.stderr + + def test_ordinary_workload_retains_root_artifacts(monkeypatch, tmp_path): state = _install_runtime(monkeypatch) result = _run(tmp_path)