Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions src/cuphoton/core/dragon.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
40 changes: 39 additions & 1 deletion tests/core/test_dragon.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
Loading