Skip to content
Open
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
8 changes: 8 additions & 0 deletions docs/reference/source/taskgraph.util.rst
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,14 @@ taskgraph.util.keyed\_by module
:undoc-members:
:show-inheritance:

taskgraph.util.memory module
----------------------------

.. automodule:: taskgraph.util.memory
:members:
:undoc-members:
:show-inheritance:

taskgraph.util.parameterization module
--------------------------------------

Expand Down
209 changes: 197 additions & 12 deletions src/taskgraph/generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,16 +6,20 @@
import inspect
import logging
import multiprocessing
import multiprocessing.connection
import os
import pickle
import platform
import signal
import sys
import traceback
from concurrent.futures import (
FIRST_COMPLETED,
ProcessPoolExecutor,
ThreadPoolExecutor,
wait,
)
from dataclasses import dataclass
from typing import Callable, Optional, Union
from typing import Callable, Optional, Union, cast

from . import filter_tasks
from .config import GraphConfig, load_graph_config
Expand All @@ -26,6 +30,7 @@
from .task import Task
from .taskgraph import TaskGraph
from .transforms.base import TransformConfig, TransformSequence
from .util.memory import gc_disabled
from .util.python_path import find_object
from .util.schema import SchemaValidationError
from .util.verify import verifications
Expand All @@ -40,6 +45,12 @@ class KindNotFound(Exception):
"""


def _max_workers():
# os.process_cpu_count is new in Python 3.13.
cpu_count = getattr(os, "process_cpu_count", os.cpu_count)
return cpu_count() or 1


@dataclass(frozen=True)
class Kind:
name: str
Expand Down Expand Up @@ -124,6 +135,16 @@ def load(cls, root_dir, graph_config, kind_name):
return cls(kind_name, path, config, graph_config)


class _RemoteTraceback(Exception):
"""The traceback of an exception raised in a child process."""

def __init__(self, tb):
self.tb = tb

def __str__(self):
return f'\n"""\n{self.tb}"""'


class TaskGraphGenerator:
"""
The central controller for taskgraph. This handles all phases of graph
Expand Down Expand Up @@ -399,6 +420,157 @@ def submit_ready_kinds():

return all_tasks

def _load_tasks_forked(self, kinds, kind_graph, parameters):
"""Load each kind in a child process forked once the kinds it depends
on are loaded.

The child shares every task loaded so far with the parent, so the
tasks of the kinds it depends on don't need to be sent to it, and it
sends its own tasks back through a pipe of its own.
"""
all_tasks = {}
tasks_by_kind = {}
# Pipe to each running child -> (kind name, pid)
running = {}
remaining = kinds.copy()
edges = set(kind_graph.edges)
max_workers = _max_workers()

try:
while remaining or running:
kinds_with_deps = {edge[0] for edge in edges}
ready_kinds = sorted(set(remaining) - kinds_with_deps)
for name in ready_kinds[: max_workers - len(running)]:
kind = remaining.pop(name)
reader, pid = self._fork_kind(
kind,
parameters,
self._get_kind_dependencies_tasks(kind, tasks_by_kind),
)
running[reader] = (name, pid)

# Waiting on nothing would block forever. _run made sure there
# is no dependency loop, so some kind is always running here.
assert running
for ready in multiprocessing.connection.wait(list(running)):
# wait returns the objects it was given, but is typed as
# also returning sockets and file descriptors.
reader = cast(multiprocessing.connection.Connection, ready)
name, pid = running[reader]
try:
result = reader.recv()
except EOFError:
# The child exited without sending anything.
result = None
del running[reader]
reader.close()
_, status = os.waitpid(pid, 0)
new_tasks = self._kind_result(name, result, status)

kind_tasks = tasks_by_kind.setdefault(name, {})
for task in new_tasks:
if task.label in all_tasks:
raise Exception("duplicate tasks with label " + task.label)
all_tasks[task.label] = task
kind_tasks[task.label] = task

edges = {e for e in edges if e[1] != name}
finally:
# Only left with running children on error.
for reader, (_, pid) in running.items():
reader.close()
os.kill(pid, signal.SIGKILL)
os.waitpid(pid, 0)

return all_tasks

def _fork_kind(self, kind, parameters, kind_dependencies_tasks):
"""Fork a child process loading the tasks of `kind`, and return the
pipe its result comes back through, and its pid."""
reader, writer = multiprocessing.Pipe(duplex=False)
# The child inherits the disabled garbage collector, so it doesn't
# write to every page it shares with us traversing the objects in them.
with gc_disabled():
pid = os.fork()
if pid == 0:
# Never leave this block, or return into the caller's code, in
# the child.
status = 1
try:
reader.close()
status = self._load_kind_in_child(
kind, parameters, kind_dependencies_tasks, writer
)
finally:
try:
sys.stdout.flush()
sys.stderr.flush()
finally:
os._exit(status)
writer.close()
return reader, pid

def _load_kind_in_child(self, kind, parameters, kind_dependencies_tasks, writer):
"""Load the tasks of `kind` and send them, or the exception loading
them raised, through `writer`. Return the child's exit status."""
try:
tasks = kind.load_tasks(
parameters, kind_dependencies_tasks, self._write_artifacts
)
except BaseException as e:
tb = traceback.format_exc()
try:
# Make sure the parent will be able to unpickle it.
pickle.loads(pickle.dumps(e))
except Exception:
e = Exception(f"{type(e).__name__}: {e}")
writer.send((e, tb, None))
return 1

try:
writer.send((None, None, tasks))
except Exception as e:
writer.send(
(
Exception(
f"Could not send the tasks of kind {kind.name} "
f"to the parent process: {type(e).__name__}: {e}"
),
traceback.format_exc(),
None,
)
)
return 1
return 0

@staticmethod
def _kind_result(name, result, status):
"""Return the tasks the child loading kind `name` sent, given what it
sent and its exit status. Raise the exception loading them raised in
the child, if any."""
if result is None:
code = os.waitstatus_to_exitcode(status)
if code < 0:
how = f"was killed by signal {signal.Signals(-code).name}"
else:
how = f"exited with status {code}"
raise Exception(
f"Process loading tasks for kind {name} {how} without sending them"
)

exc, tb, tasks = result
if exc is not None:
exc.__cause__ = _RemoteTraceback(tb)
if isinstance(exc, SchemaValidationError):
logger.error(f"Error loading tasks for kind {name}:\n{exc}")
else:
logger.error(
f"Error loading tasks for kind {name}:",
exc_info=(type(exc), exc, exc.__traceback__),
)
raise exc
return tasks

def _run(self):
logger.info("Loading graph configuration.")
graph_config = load_graph_config(self.root_dir)
Expand Down Expand Up @@ -446,21 +618,32 @@ def _run(self):
for kind in kinds.values():
for dep in kind.config.get("kind-dependencies", []):
edges.add((kind.name, dep, "kind-dependency"))
for _, dep, _ in edges:
if dep not in kinds:
message = f'Could not find the kind "{dep}"\nAvailable kinds:\n'
for k in sorted(kinds):
message += f' - "{k}"\n'
raise Exception(message)
kind_graph = Graph(frozenset(kinds), frozenset(edges))

if target_kinds:
kind_graph = kind_graph.transitive_closure(
set(target_kinds) | {"docker-image"}
)

# The task loaders wait for each kind's dependencies to be loaded
# before loading it, so a dependency loop would hang them or silently
# leave kinds out. This raises if there is one.
kind_graph.visit_postorder()

yield "kind_graph", kind_graph

logger.info("Generating full task set")

# The next block deals with enabling parallel kind processing, which
# currently has different support on different platforms. In summary:
# * Parallel kind processing is supported and enabled by default on
# Linux. We use multiple processes by default, but experimental
# Linux. We fork a process per kind by default, but experimental
# support for multiple threads can be enabled instead.
# * On other platforms, we have experimental support for parallel
# kind processing with multiple threads.
Expand All @@ -480,21 +663,19 @@ def load_tasks():
if os.environ.get("TASKGRAPH_SERIAL"):
return self._load_tasks_serial(kinds, kind_graph, parameters)
elif os.environ.get("TASKGRAPH_USE_THREADS"):
executor = ThreadPoolExecutor(max_workers=os.process_cpu_count())
else:
executor = ProcessPoolExecutor(
mp_context=multiprocessing.get_context("fork")
executor = ThreadPoolExecutor(max_workers=_max_workers())
return self._load_tasks_parallel(
kinds, kind_graph, parameters, executor
)
return self._load_tasks_parallel(
kinds, kind_graph, parameters, executor
)
else:
return self._load_tasks_forked(kinds, kind_graph, parameters)
else:
if os.environ.get("TASKGRAPH_SERIAL") or not os.environ.get(
"TASKGRAPH_USE_THREADS"
):
return self._load_tasks_serial(kinds, kind_graph, parameters)
else:
executor = ThreadPoolExecutor(max_workers=os.process_cpu_count())
executor = ThreadPoolExecutor(max_workers=_max_workers())
return self._load_tasks_parallel(
kinds, kind_graph, parameters, executor
)
Expand Down Expand Up @@ -604,7 +785,11 @@ def load_tasks():
def _run_until(self, name):
while name not in self._run_results:
try:
k, v = next(self._run) # type: ignore
# Generation creates millions of objects that stay alive, so
# collecting garbage meanwhile only traverses them for nothing.
# The callers' code runs with the collector in its usual state.
with gc_disabled():
k, v = next(self._run) # type: ignore
except StopIteration:
raise AttributeError(f"No such run result {name}")
self._run_results[k] = v
Expand Down
18 changes: 10 additions & 8 deletions src/taskgraph/taskgraph.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@

from .graph import Graph
from .task import Task
from .util.memory import gc_disabled


@dataclass(frozen=True)
Expand Down Expand Up @@ -58,13 +59,14 @@ def from_json(cls, tasks_dict):
"""
tasks = {}
edges = set()
for key, value in tasks_dict.items():
tasks[key] = Task.from_json(value)
if "task_id" in value:
tasks[key].task_id = value["task_id"]
for depname, dep in value["dependencies"].items():
# Task filtering can cause dependencies to be removed from the graph.
if dep in tasks_dict:
edges.add((key, dep, depname))
with gc_disabled():
for key, value in tasks_dict.items():
tasks[key] = Task.from_json(value)
if "task_id" in value:
tasks[key].task_id = value["task_id"]
for depname, dep in value["dependencies"].items():
# Task filtering can cause dependencies to be removed from the graph.
if dep in tasks_dict:
edges.add((key, dep, depname))
task_graph = cls(tasks, Graph(frozenset(tasks), edges))
return tasks, task_graph
24 changes: 24 additions & 0 deletions src/taskgraph/util/memory.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
# This Source Code Form is subject to the terms of the Mozilla Public
# License, v. 2.0. If a copy of the MPL was not distributed with this
# file, You can obtain one at http://mozilla.org/MPL/2.0/.

import gc
from contextlib import contextmanager


@contextmanager
def gc_disabled():
"""Disable the garbage collector for the duration of the block, restoring
its previous state afterwards.

Generating or loading a task graph creates millions of objects that stay
alive, so collections keep traversing a growing heap without freeing
anything.
"""
was_enabled = gc.isenabled()
gc.disable()
try:
yield
finally:
if was_enabled:
gc.enable()
Loading
Loading