diff --git a/.gitignore b/.gitignore index c53f216a3..dad032970 100644 --- a/.gitignore +++ b/.gitignore @@ -1,2 +1,2 @@ .clang-format -.mypy-cache +.mypy_cache diff --git a/python/CMakeLists.txt b/python/CMakeLists.txt index c73e58f72..7d37c12c0 100644 --- a/python/CMakeLists.txt +++ b/python/CMakeLists.txt @@ -154,7 +154,8 @@ python_module( MODULE_INIT revng/internal/__init__.py MODULE_FILES - revng/internal/py.typed) + revng/internal/py.typed + revng/internal/pipebox.py) # # Install revng.model (including autogenerated classes) @@ -230,7 +231,8 @@ python_module(TARGET_NAME revng-dump-model WHEEL revng_internal MODULE_FILES # set(REVNG_CLI_MODULE_FILES revng/internal/cli/commands_registry.py revng/internal/cli/__init__.py - revng/internal/cli/revng.py revng/internal/cli/support.py) + revng/internal/cli/revng.py revng/internal/cli/revng2.py + revng/internal/cli/support.py) python_module(TARGET_NAME revng-python-cli WHEEL revng_internal MODULE_FILES ${REVNG_CLI_MODULE_FILES}) @@ -374,9 +376,61 @@ python_module( MODULE_GENERATED_FILES revng/pipeline_description/_generated.py) +# +# Install pypeline +# +set(REVNG_PYPELINE_MODULE_FILES + revng/pypeline/cli/pipeline/__init__.py + revng/pypeline/cli/pipeline/run_analysis.py + revng/pypeline/cli/pipeline/run_pipe.py + revng/pypeline/cli/project/__init__.py + revng/pypeline/cli/project/analyze.py + revng/pypeline/cli/project/artifact.py + revng/pypeline/cli/__init__.py + revng/pypeline/cli/utils.py + revng/pypeline/analysis.py + revng/pypeline/container.py + revng/pypeline/graph.py + revng/pypeline/main.py + revng/pypeline/model.py + revng/pypeline/object.py + revng/pypeline/pipeline_node.py + revng/pypeline/pipeline_parser.py + revng/pypeline/pipeline.py + revng/pypeline/pipeline_schema.yml + revng/pypeline/schedule/__init__.py + revng/pypeline/schedule/scheduled_task.py + revng/pypeline/schedule/schedule.py + revng/pypeline/storage/memory.py + revng/pypeline/storage/null.py + revng/pypeline/storage/sqlite3.py + revng/pypeline/storage/storage_provider.py + revng/pypeline/storage/util.py + revng/pypeline/storage/__init__.py + revng/pypeline/task/__init__.py + revng/pypeline/task/pipe.py + revng/pypeline/task/requests.py + revng/pypeline/task/savepoint.py + revng/pypeline/task/task.py + revng/pypeline/utils/default_dict_from_key.py + revng/pypeline/utils/registry.py + revng/pypeline/utils/cabc.py + revng/pypeline/utils/__init__.py) +python_module( + TARGET_NAME + revng-python-pypeline + WHEEL + pypeline + MODULE_INIT + revng/pypeline/__init__.py + MODULE_FILES + ${REVNG_PYPELINE_MODULE_FILES}) + # # Generate actual wheel rules # python_wheel(NAME revng PYPROJECT revng/pyproject.toml README revng/README.md) python_wheel(NAME revng_internal PYPROJECT revng/internal/pyproject.toml README revng/internal/README.md) +python_wheel(NAME pypeline PYPROJECT revng/pypeline/pyproject.toml README + revng/pypeline/README.md) diff --git a/python/revng/internal/cli/revng2.py b/python/revng/internal/cli/revng2.py new file mode 100644 index 000000000..8db86ae3f --- /dev/null +++ b/python/revng/internal/cli/revng2.py @@ -0,0 +1,53 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +""" +This is just a wrapper over `pype` that sets pipebox to the revng pipebox path. +The path is computed relatively to this file, so this should work regardless of +where revng is installed. +""" + +import logging +import os +import sys +from pathlib import Path +from typing import Sequence + +import click + +from revng.pypeline.cli.utils import LazyGroup +from revng.pypeline.main import import_pipebox + +logger = logging.getLogger("revng2") +logger.setLevel(logging.DEBUG) +logger.addHandler(logging.StreamHandler(sys.stderr)) + + +@click.group( + cls=LazyGroup, + lazy_subcommands={ + "pipeline": "revng.pypeline.cli.pipeline:pipeline", + "project": "revng.pypeline.cli.project:project", + }, +) +def cli(): + pass + + +def main(args: Sequence[str]) -> None: + # This should resolve to the full path of revng/internal/pipebox.py + pipebox_path = Path(__file__).parent.parent / "pipebox.py" + import_pipebox(str(pipebox_path), "_REVNG2_COMPLETE" in os.environ) + # pylint: disable=E1120 no-value-for-parameter + # This is ok as click will pass the pipebox argument automatically + cli(args=args) + + +def run(): + """Run the pipeline from the command line using the shell environment.""" + main(sys.argv[1:]) + + +if __name__ == "__main__": + run() diff --git a/python/revng/internal/pipebox.py b/python/revng/internal/pipebox.py new file mode 100644 index 000000000..ac19c01eb --- /dev/null +++ b/python/revng/internal/pipebox.py @@ -0,0 +1,8 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +""" +Placeholder pipebox file that will be filled with the new nanobind-based +implementations +""" diff --git a/python/revng/internal/pyproject.toml b/python/revng/internal/pyproject.toml index e90baa7b0..6b0246fa6 100644 --- a/python/revng/internal/pyproject.toml +++ b/python/revng/internal/pyproject.toml @@ -50,6 +50,7 @@ revng = [ [project.entry-points.console_scripts] revng = "revng.internal.cli.revng:main" +revng2 = "revng.internal.cli.revng2:run" [tool.setuptools.dynamic] version = {attr = "revng.internal.__version__"} diff --git a/python/revng/pypeline/README.md b/python/revng/pypeline/README.md new file mode 100644 index 000000000..8b1378917 --- /dev/null +++ b/python/revng/pypeline/README.md @@ -0,0 +1 @@ + diff --git a/python/revng/pypeline/__init__.py b/python/revng/pypeline/__init__.py new file mode 100644 index 000000000..6ffdda6b1 --- /dev/null +++ b/python/revng/pypeline/__init__.py @@ -0,0 +1,35 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +# It's important that the `cli` submodule is NOT imported here +# as for the cli to work we need to first load the pipebox and then +# load the cli submodules so they can access the custom types + +__version__ = "@VERSION@" + +from .analysis import Analysis +from .container import Container +from .model import Model +from .object import Kind, ObjectID +from .task.pipe import Pipe +from .utils.registry import get_singleton, register_all_subclasses + + +def initialize_pypeline() -> None: + """ + This function is used to initialize the pypeline module. + It is has to be called just after importing the pipebox. + + All of this can be done more robustly using `__init_subclass__`, + but it's not well supported by nanobind, so we are forced to do it + manually here. + """ + register_all_subclasses(Analysis) + register_all_subclasses(Pipe) + register_all_subclasses(Container) + register_all_subclasses(Model, singleton=True) + register_all_subclasses(Kind, singleton=True) + register_all_subclasses(ObjectID, singleton=True) + kind_ty = get_singleton(Kind) # type: ignore[type-abstract] + kind_ty._init_type() diff --git a/python/revng/pypeline/analysis.py b/python/revng/pypeline/analysis.py new file mode 100644 index 000000000..2221f08dd --- /dev/null +++ b/python/revng/pypeline/analysis.py @@ -0,0 +1,65 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from dataclasses import dataclass + +from .container import Container, ContainerDeclaration +from .model import Model +from .object import ObjectSet +from .pipeline_node import PipelineNode +from .utils.cabc import ABC, abstractmethod + + +class Analysis(ABC): + """ + An analysis makes changes to the model. In order to do this, it might inspect + previously produced results of the pipeline. + + An analysis is the way in which users are expected to make changes to the model. + The changes applied by a run of an analysis might lead to invalidate certain + objects in save points. + """ + + __slots__: tuple = ("name",) + + def __init__(self, name: str): + """ + Initialize the analysis with a name. + The name is used for debugging and logging purposes. + """ + self.name = name + + @classmethod + @abstractmethod + def signature(cls) -> tuple[type[Container], ...]: + """ + The containers required by the analysis, it needs to be a class property + to auto-generate the CLI commands. + """ + raise NotImplementedError() + + @abstractmethod + def run( + self, + model: Model, + containers: list[Container], + incoming: list[ObjectSet], + configuration: str, + ): + """ + Run the analysis on the model, using the containers and + incoming requests. The analysis will modify inplace the + model so you must make a copy before running it, so you + can compute the diff for invalidation purposes. + """ + raise NotImplementedError() + + +@dataclass(frozen=True, slots=True) +class AnalysisBinding: + """Allows to bind an analysis to a pipeline node.""" + + analysis: Analysis + bindings: tuple[ContainerDeclaration, ...] + node: PipelineNode diff --git a/python/revng/pypeline/cli/__init__.py b/python/revng/pypeline/cli/__init__.py new file mode 100644 index 000000000..37804b081 --- /dev/null +++ b/python/revng/pypeline/cli/__init__.py @@ -0,0 +1,5 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +# This is deliberately empty for lazy loading of subcommands diff --git a/python/revng/pypeline/cli/pipeline/__init__.py b/python/revng/pypeline/cli/pipeline/__init__.py new file mode 100644 index 000000000..0094c5110 --- /dev/null +++ b/python/revng/pypeline/cli/pipeline/__init__.py @@ -0,0 +1,19 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +import click + +from revng.pypeline.cli.utils import LazyGroup + + +@click.group( + cls=LazyGroup, + lazy_subcommands={ + "run_pipe": "revng.pypeline.cli.pipeline.run_pipe:run_pipe", + "run_analysis": "revng.pypeline.cli.pipeline.run_analysis:run_analysis", + }, + help="Low-level pipeline commands (plumbing)", +) +def pipeline() -> None: + pass diff --git a/python/revng/pypeline/cli/pipeline/run_analysis.py b/python/revng/pypeline/cli/pipeline/run_analysis.py new file mode 100644 index 000000000..787d17c3c --- /dev/null +++ b/python/revng/pypeline/cli/pipeline/run_analysis.py @@ -0,0 +1,184 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +import logging +import sys + +import click + +from revng.pypeline.analysis import Analysis +from revng.pypeline.cli.utils import build_arg_objects, build_help_text, compute_objects +from revng.pypeline.cli.utils import normalize_whitespace +from revng.pypeline.container import ContainerDeclaration, load_container +from revng.pypeline.model import Model, ReadOnlyModel +from revng.pypeline.object import ObjectSet +from revng.pypeline.task.task import TaskArgument, TaskArgumentAccess +from revng.pypeline.utils.registry import get_registry, get_singleton + +logger = logging.getLogger(__name__) + + +class RunAnalysisGroup(click.Group): + """We need to create a custom command for each analysis we loaded from the registry. + Since we already have to generate the code dynamically, we do it lazily so + we generate only the commands that are requested. + This is based on the LazyGroup from revng.pypeline.cli.utils.""" + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.registry: dict[str, type[Analysis]] = get_registry(Analysis) + self.model_ty: type[Model] = get_singleton(Model) # type: ignore[type-abstract] + + def list_commands(self, ctx): + base = super().list_commands(ctx) + return base + sorted(self.registry.keys()) + + def get_command(self, ctx, cmd_name): + if cmd_name in self.registry: + return self._build_analysis_command(cmd_name) + return super().get_command(ctx, cmd_name) + + def _build_analysis_command(self, analysis_name: str): + """Dynamically create a command for running an analysis.""" + analysis_ty: type[Analysis] = self.registry[analysis_name] + + if analysis_ty.__doc__: + help_text = click.wrap_text(f"\n{normalize_whitespace(analysis_ty.__doc__)}") + else: + help_text = f"Run the analysis: {analysis_name}" + + help_text = build_help_text( + prologue=help_text, + args=[ + TaskArgument( + name=container_type.__name__, + container_type=container_type, + access=TaskArgumentAccess.READ, + help_text=normalize_whitespace(container_type.__doc__ or ""), + ) + for container_type in analysis_ty.signature() + ], + ) + + # Build the actual function that will be the command + run_analysis_command = build_run_analysis_command( + analysis_name=analysis_name, + help_text=help_text, + analysis_ty=analysis_ty, + model_ty=self.model_ty, + ) + + config = getattr( + analysis_ty, "configuration_help", f"Configuration for the analysis '{analysis_name}'." + ) + if config is not None: + run_analysis_command = click.option( + "-c", + "--configuration", + type=str, + default="", + help=normalize_whitespace(config), + )(run_analysis_command) + + # For each argument, call the `click.argument` decorator to dynamically add + # them to the command + for arg in analysis_ty.signature(): + run_analysis_command = click.argument( + arg.__name__, + type=click.Path(exists=True, dir_okay=False, readable=True), + )(run_analysis_command) + run_analysis_command = build_arg_objects( + ContainerDeclaration( + name=arg.__name__, + container_type=arg, + ) + )(run_analysis_command) + + return run_analysis_command + + +def build_run_analysis_command( + analysis_name: str, + help_text: str, + analysis_ty: type[Analysis], + model_ty: type[Model], +): + @click.command(name=analysis_name, help=help_text) + @click.argument( + "model", + type=click.Path(exists=True, dir_okay=False, readable=True), + required=True, + ) + @click.option( + "--list", + type=bool, + is_flag=True, + default=False, + help="List the available objects for each argument.", + ) + def run_analysis_command( + model: str, + configuration: str, + **kwargs, + ) -> None: + logger.debug("Running analysis: `%s`", analysis_name) + logger.debug("configuration: `%s`", configuration) + logger.debug("model: `%s`", model) + logger.debug("and kwargs: `%s`", kwargs) + + analysis = analysis_ty( + name=analysis_name, + ) + # Load the model + loaded_model: Model = model_ty() + with open(model, "rb") as model_file: + loaded_model.deserialize(model_file.read()) + + logger.debug("Model loaded: `%s`", loaded_model) + + # Load the containers with args form the command line + containers = [] + for arg in analysis.signature(): + arg_name = arg.__name__ + path = kwargs[arg_name] + container = load_container(arg, path) + logger.debug( + "Loaded container from `%s` for argument `%s`: `%r`", path, arg_name, container + ) + containers.append(container) + + # Compute the requests for the incoming containers of the + # analysis + incoming: list[ObjectSet] = [] + for arg in analysis.signature(): + # If the argument is writable, we need to request the objects + incoming.append( + compute_objects( + model=ReadOnlyModel(loaded_model), + arg_name=arg.__name__, + kind=arg.kind, + kwargs=kwargs, + ) + ) + + # Finally, run the analysis + analysis.run( + model=loaded_model, + containers=containers, + incoming=incoming, + configuration=configuration, + ) + logger.debug("Analysis run completed") + # Print on stdout the raw bytes of the modified model + sys.stdout.buffer.write(loaded_model.serialize()) + + return run_analysis_command + + +@click.group( + cls=RunAnalysisGroup, + help="Run an analysis", +) +def run_analysis() -> None: + pass diff --git a/python/revng/pypeline/cli/pipeline/run_pipe.py b/python/revng/pypeline/cli/pipeline/run_pipe.py new file mode 100644 index 000000000..dd192798a --- /dev/null +++ b/python/revng/pypeline/cli/pipeline/run_pipe.py @@ -0,0 +1,245 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +import logging + +import click + +from revng.pypeline.cli.utils import build_arg_objects, build_help_text, compute_objects +from revng.pypeline.cli.utils import normalize_whitespace +from revng.pypeline.container import dump_container, load_container +from revng.pypeline.model import Model, ReadOnlyModel +from revng.pypeline.object import ObjectSet +from revng.pypeline.task.pipe import Pipe +from revng.pypeline.task.task import TaskArgumentAccess +from revng.pypeline.utils.registry import get_registry, get_singleton + +logger = logging.getLogger(__name__) + + +class RunPipeGroup(click.Group): + """We need to create a custom command for each pipe we loaded from the registry. + Since we already have to generate the code dynamically, we do it lazily so + we generate only the commands that are requested. + This is based on the LazyGroup from revng.pypeline.cli.utils.""" + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.registry: dict[str, type[Pipe]] = get_registry(Pipe) + self.model_ty: type[Model] = get_singleton(Model) # type: ignore[type-abstract] + + def list_commands(self, ctx): + base = super().list_commands(ctx) + return base + sorted(self.registry.keys()) + + def get_command(self, ctx, cmd_name): + if cmd_name in self.registry: + return self._build_pipe_command(cmd_name) + return super().get_command(ctx, cmd_name) + + def _build_pipe_command(self, pipe_name: str): + """Dynamically create a command for running a pipe.""" + pipe_ty: type[Pipe] = self.registry[pipe_name] + + if pipe_ty.__doc__: + help_text = click.wrap_text(f"\n{normalize_whitespace(pipe_ty.__doc__)}") + else: + help_text = f"Run the pipe: {pipe_name}" + + help_text = build_help_text( + prologue=help_text, + args=pipe_ty.signature(), + ) + + # Add options for static configuration and configuration, only if the + # pipe doesn't disable them by defining them as None + static_config = ( + pipe_ty.static_configuration_help() + or f"Static configuration for the pipe '{pipe_name}'." + ) + + # Build the actual function that will be the command + run_pipe_command = build_pipe_command( + pipe_name=pipe_name, + help_text=help_text, + pipe_ty=pipe_ty, + model_ty=self.model_ty, + ) + + # Decorate it to add the arguments it needs + if static_config is not None: + run_pipe_command = click.option( + "-s", + "--static-configuration", + type=str, + default="", + help=normalize_whitespace(static_config), + )(run_pipe_command) + config = getattr( + pipe_ty, "configuration_help", f"Configuration for the pipe '{pipe_name}'." + ) + if config is not None: + run_pipe_command = click.option( + "-c", + "--configuration", + type=str, + default="", + help=normalize_whitespace(config), + )(run_pipe_command) + + # For each argument, call the `click.argument` decorator to dynamically add + # them to the command + for arg in pipe_ty.signature(): + if arg.access == TaskArgumentAccess.READ_WRITE: + run_pipe_command = click.argument( + f"{arg.name}-input", + type=click.Path(exists=True, dir_okay=False, readable=True), + )(run_pipe_command) + run_pipe_command = click.argument( + f"{arg.name}-output", + type=click.Path(dir_okay=False, writable=True), + )(run_pipe_command) + elif arg.access == TaskArgumentAccess.READ: + run_pipe_command = click.argument( + arg.name, + type=click.Path(exists=True, dir_okay=False, readable=True), + )(run_pipe_command) + elif arg.access == TaskArgumentAccess.WRITE: + run_pipe_command = click.argument( + arg.name, + type=click.Path(dir_okay=False, writable=True), + )(run_pipe_command) + else: + raise ValueError("Unreachable code: unknown access type %s") + # For output arguments, we can specify which objects we want + # to request + if arg.access & TaskArgumentAccess.WRITE: + run_pipe_command = build_arg_objects(arg)(run_pipe_command) + + return run_pipe_command + + +def build_pipe_command( + pipe_name: str, + help_text: str, + pipe_ty: type[Pipe], + model_ty: type[Model], +): + @click.command(name=pipe_name, help=help_text) + @click.argument( + "model", + type=click.Path(exists=True, dir_okay=False, readable=True), + required=True, + ) + @click.option( + "--list", + type=bool, + is_flag=True, + default=False, + help="List the available objects for each argument.", + ) + def run_pipe_command( + model: str, static_configuration: str, configuration: str, **kwargs + ) -> None: + logger.debug("Running pipe: %s", pipe_name) + logger.debug("with static configuration: %s", static_configuration) + logger.debug("configuration: %s", configuration) + logger.debug("model: %s", model) + logger.debug("and kwargs: %s", kwargs) + + # Create the pipe + pipe = pipe_ty( + name=pipe_name, + static_configuration=static_configuration, + ) + # Load the model + loaded_model: Model = model_ty() + with open(model, "rb") as model_file: + loaded_model.deserialize(model_file.read()) + # Load the containers with args form the command line + containers = [] + for arg in pipe.arguments: + # Write-only containers can be empty + if arg.access == TaskArgumentAccess.WRITE: + containers.append(arg.container_type()) + continue + + # Otherwise we need to load the container from the filesystem + if arg.access == TaskArgumentAccess.READ_WRITE: + path = kwargs[f"{arg.name}-input"] + else: + path = kwargs[arg.name] + containers.append( + load_container( + arg.container_type, + path, + ) + ) + + # From the command line figure out the requests for each argument + outgoing: list[ObjectSet] = [] + for arg in pipe.signature(): + if arg.access == TaskArgumentAccess.READ: + # If the argument is read-only, we don't need to request anything + continue + # If the argument is writable, we need to request the objects + outgoing.append( + compute_objects( + model=ReadOnlyModel(loaded_model), + arg_name=arg.name, + kind=arg.container_type.kind, + kwargs=kwargs, + ) + ) + + # Ask the pipe for the requests it needs + incoming = pipe.prerequisites_for( + model=ReadOnlyModel(loaded_model), + requests=outgoing, + ) + # Ensure that the user provided all the required arguments + for container, request in zip(containers, incoming): + if not request: + # If the request is empty, we don't need to load anything + continue + if container.kind not in request: + raise click.UsageError( + f"Container {container} does not have the required objects: {request}" + ) + + # Finally, run the pipe + object_deps = pipe.run( + model=ReadOnlyModel(loaded_model), + containers=containers, + incoming=incoming, + outgoing=outgoing, + configuration=configuration, + ) + logger.debug("Pipe run completed, object dependencies: `%s`", object_deps) + + # Dump back the modified containers to the filesystem + for arg, container in zip(pipe.signature(), containers): + if arg.access == TaskArgumentAccess.READ: + continue + # If the argument is writable, we dump the container + # to the filesystem + if arg.access == TaskArgumentAccess.READ_WRITE: + path = kwargs[f"{arg.name}-output"] + else: + path = kwargs[arg.name] + logger.info("Dumping container %s to %s", arg.name, path) + dump_container( + container, + path, + ) + + return run_pipe_command + + +@click.group( + cls=RunPipeGroup, + help="Run a pipe", +) +def run_pipe() -> None: + pass diff --git a/python/revng/pypeline/cli/project/__init__.py b/python/revng/pypeline/cli/project/__init__.py new file mode 100644 index 000000000..eeb5a2469 --- /dev/null +++ b/python/revng/pypeline/cli/project/__init__.py @@ -0,0 +1,35 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +import os + +import click + +from revng.pypeline.cli.utils import EagerParsedPath, LazyGroup +from revng.pypeline.pipeline_parser import load_pipeline_yaml_file + + +@click.group( + cls=LazyGroup, + lazy_subcommands={ + "analyze": "revng.pypeline.cli.project.analyze:analyze", + "artifact": "revng.pypeline.cli.project.artifact:get_artifact", + }, + help="Project commands (porcelain)", +) +@click.option( + "--pipeline", + type=EagerParsedPath( + name="pipeline", + parser=load_pipeline_yaml_file, + ), + help=( + "Path to the pipeline file. Defaults to the `PIPELINE` environment " + "variable, then 'pipeline.yml'." + ), + default=os.environ.get("PIPELINE", "pipeline.yml"), + expose_value=False, +) +def project() -> None: + pass diff --git a/python/revng/pypeline/cli/project/analyze.py b/python/revng/pypeline/cli/project/analyze.py new file mode 100644 index 000000000..1f439d555 --- /dev/null +++ b/python/revng/pypeline/cli/project/analyze.py @@ -0,0 +1,170 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +import logging +import sys + +import click + +from revng.pypeline.cli.utils import build_arg_objects, build_help_text, compute_objects +from revng.pypeline.cli.utils import list_objects_for_container, normalize_whitespace +from revng.pypeline.cli.utils import storage_provider_factory +from revng.pypeline.model import Model, ReadOnlyModel +from revng.pypeline.pipeline import AnalysisBinding, Pipeline +from revng.pypeline.task.requests import Requests +from revng.pypeline.utils.registry import get_singleton + +logger = logging.getLogger(__name__) + + +class AnalyzeGroup(click.Group): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.model_ty: type[Model] = get_singleton(Model) # type: ignore[type-abstract] + + def list_commands(self, ctx): + base = super().list_commands(ctx) + pipeline = ctx.obj.get("pipeline") + if pipeline is None: + return base + return base + sorted(pipeline.analyses.keys()) + + def get_command(self, ctx, cmd_name): + pipeline = ctx.obj.get("pipeline") + if pipeline is None: + return super().get_command(ctx, cmd_name) + if cmd_name not in pipeline.analyses: + return super().get_command(ctx, cmd_name) + return self._build_analysis_command( + analysis_name=cmd_name, + pipeline=pipeline, + ) + + def _build_analysis_command(self, analysis_name: str, pipeline: Pipeline): + """Dynamically create a command for running an analysis.""" + analysis_binding: AnalysisBinding = pipeline.analyses[analysis_name] + + if analysis_binding.analysis.__doc__: + help_text = click.wrap_text( + f"\n{normalize_whitespace(analysis_binding.analysis.__doc__)}" + ) + else: + help_text = f"Run the analysis: {analysis_name}" + + help_text = build_help_text(prologue=help_text, args=[]) + + # Build the actual function that will be the command + run_analysis_command = build_analysis_command( + analysis_binding=analysis_binding, + help_text=help_text, + model_ty=self.model_ty, + pipeline=pipeline, + ) + + config = getattr( + analysis_binding.analysis, + "configuration_help", + f"Configuration for the analysis '{analysis_name}'.", + ) + if config is not None: + run_analysis_command = click.option( + "-c", + "--configuration", + type=str, + default="", + help=normalize_whitespace(config), + )(run_analysis_command) + + # For each argument, call the `click.argument` decorator to dynamically add + # them to the command + for container_decl in analysis_binding.bindings: + run_analysis_command = build_arg_objects(container_decl)(run_analysis_command) + + return run_analysis_command + + +def build_analysis_command( + analysis_binding: AnalysisBinding, + help_text: str, + model_ty: type[Model], + pipeline: Pipeline, +): + analysis_name: str = analysis_binding.analysis.name + + @click.command(name=analysis_name, help=help_text) + @click.argument( + "model", + type=click.Path(exists=True, dir_okay=False, readable=True), + required=True, + ) + @click.option( + "--list", + type=bool, + is_flag=True, + default=False, + help="List the available objects for each argument.", + ) + def run_analysis_command( + model: str, + configuration: str, + **kwargs, + ) -> None: + logger.debug("Running analysis: `%s`", analysis_name) + logger.debug("configuration: `%s`", configuration) + logger.debug("model: `%s`", model) + logger.debug("and kwargs: `%s`", kwargs) + + # Load the model + storage_provider = storage_provider_factory(model_path=model) + loaded_model: Model = model_ty() + loaded_model.deserialize(storage_provider.get_model()) + + logger.debug("Model loaded: `%s`", loaded_model) + + if kwargs["list"]: + # If the user requested to list the available objects, we print them + # and exit + for container_decl in analysis_binding.bindings: + list_objects_for_container( + model=ReadOnlyModel(loaded_model), + arg_name=container_decl.name, + kind=container_decl.container_type.kind, + ) + # Space between containers + print() + return + + # Compute the requests for the incoming containers of the + # analysis + incoming = Requests() + for container_decl in analysis_binding.bindings: + incoming[container_decl] = compute_objects( + model=ReadOnlyModel(loaded_model), + arg_name=container_decl.name, + kind=container_decl.container_type.kind, + kwargs=kwargs, + ) + + # Finally, run the analysis + new_model = pipeline.run_analysis( + model=ReadOnlyModel(loaded_model), + analysis_name=analysis_name, + requests=incoming, + analysis_configuration=configuration, + pipeline_configuration={}, + storage_provider=storage_provider, + ) + logger.debug("Analysis run completed") + # Print on stdout the raw bytes of the modified model + sys.stdout.buffer.write(new_model.serialize()) + + return run_analysis_command + + +@click.group( + cls=AnalyzeGroup, + help="Run an analysis", +) +def analyze() -> None: + pass diff --git a/python/revng/pypeline/cli/project/artifact.py b/python/revng/pypeline/cli/project/artifact.py new file mode 100644 index 000000000..6def537f7 --- /dev/null +++ b/python/revng/pypeline/cli/project/artifact.py @@ -0,0 +1,184 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +import logging + +import click + +from revng.pypeline.cli.utils import build_arg_objects, build_help_text, compute_objects +from revng.pypeline.cli.utils import normalize_whitespace, storage_provider_factory +from revng.pypeline.container import dump_container +from revng.pypeline.model import Model, ReadOnlyModel +from revng.pypeline.object import ObjectSet +from revng.pypeline.pipeline import Artifact, Pipeline +from revng.pypeline.task.task import TaskArgument, TaskArgumentAccess +from revng.pypeline.utils.registry import get_singleton + +logger = logging.getLogger(__name__) + + +class ArtifactGroup(click.Group): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.model_ty: type[Model] = get_singleton(Model) # type: ignore[type-abstract] + + def list_commands(self, ctx): + base = super().list_commands(ctx) + pipeline = ctx.obj.get("pipeline") + if pipeline is None: + return base + return base + sorted(pipeline.artifacts.keys()) + + def get_command(self, ctx, cmd_name): + pipeline = ctx.obj.get("pipeline") + if pipeline is None: + return super().get_command(ctx, cmd_name) + if cmd_name not in pipeline.artifacts: + return super().get_command(ctx, cmd_name) + return self._build_artifact_command( + artifact_name=cmd_name, + pipeline=pipeline, + ) + + def _build_artifact_command(self, artifact_name: str, pipeline: Pipeline): + """Dynamically create a command for getting an artifact.""" + artifact: Artifact = pipeline.artifacts[artifact_name] + + if artifact.__doc__: + help_text = click.wrap_text(f"\n{normalize_whitespace(artifact.__doc__)}") + else: + help_text = f"Get the artifact: {artifact_name}" + + help_text = build_help_text( + prologue=help_text, + args=[ + TaskArgument( + name=artifact.container.name, + container_type=artifact.container.container_type, + access=TaskArgumentAccess.WRITE, + help_text=artifact.container.container_type.__doc__ or "", + ) + ], + ) + + # Build the actual function that will be the command + run_artifact_command = build_artifact_command( + artifact=artifact, + help_text=help_text, + model_ty=self.model_ty, + pipeline=pipeline, + ) + + config = getattr( + artifact, + "configuration_help", + f"Configuration for the artifact '{artifact_name}'.", + ) + if config is not None: + run_artifact_command = click.option( + "-c", + "--configuration", + type=str, + default="", + help=normalize_whitespace(config), + )(run_artifact_command) + + # For the only container, call the `click.argument` decorator to + # dynamically add its `objects` argument to the command + run_artifact_command = click.argument( + artifact.container.name, + type=click.Path(dir_okay=False, writable=True), + )(run_artifact_command) + run_artifact_command = build_arg_objects(artifact.container)(run_artifact_command) + + return run_artifact_command + + +def build_artifact_command( + artifact: Artifact, + help_text: str, + model_ty: type[Model], + pipeline: Pipeline, +): + artifact_name: str = artifact.name + + @click.command(name=artifact_name, help=help_text) + @click.argument( + "model", + type=click.Path(exists=True, dir_okay=False, readable=True), + required=True, + ) + @click.option( + "--list", + type=bool, + is_flag=True, + default=False, + help="List the available objects for each argument.", + ) + def run_analysis_command( + model: str, + configuration: str, + **kwargs, + ) -> None: + logger.debug("Running artifact: `%s`", artifact_name) + logger.debug("configuration: `%s`", configuration) + logger.debug("model: `%s`", model) + logger.debug("and kwargs: `%s`", kwargs) + + # Load the model + storage_provider = storage_provider_factory(model_path=model) + loaded_model: Model = model_ty() + loaded_model.deserialize(storage_provider.get_model()) + + logger.debug("Model loaded: `%s`", loaded_model) + + arg_name = artifact.container.name + artifact_kind = artifact.container.container_type.kind + if kwargs["list"]: + # If the user requested to list the available objects, we print them + # and exit + print(f"Available objects for `{arg_name}` kind: {artifact_kind.__name__}") + for obj in loaded_model.all_objects(artifact_kind): + print(f" - {obj}") + return + + # Compute the requests for the incoming containers of the + # analysis + incoming: ObjectSet = loaded_model.all_objects( + artifact_kind, + ) + # If the argument is writable, we need to request the objects + incoming = compute_objects( + model=ReadOnlyModel(loaded_model), + arg_name=arg_name, + kind=artifact_kind, + kwargs=kwargs, + ) + + # Finally, run the analysis + res_container = pipeline.get_artifact( + model=ReadOnlyModel(loaded_model), + artifact=artifact, + requests=incoming, + pipeline_configuration={}, + storage_provider=storage_provider, + ) + logger.debug("Artifact computed") + + res_path = kwargs[arg_name] + logger.debug("Writing result to: `%s`", res_path) + dump_container( + res_container, + res_path, + ) + + return run_analysis_command + + +@click.group( + cls=ArtifactGroup, + help="Compute an Artifact", +) +def get_artifact() -> None: + pass diff --git a/python/revng/pypeline/cli/utils.py b/python/revng/pypeline/cli/utils.py new file mode 100644 index 000000000..8e3706efe --- /dev/null +++ b/python/revng/pypeline/cli/utils.py @@ -0,0 +1,283 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +import importlib +import importlib.util +import logging +import os +import re +from collections.abc import Sequence +from typing import Any, Callable, Optional + +import click + +from revng.pypeline.container import ContainerDeclaration +from revng.pypeline.model import ReadOnlyModel +from revng.pypeline.object import Kind, ObjectID, ObjectSet +from revng.pypeline.storage import storage_provider_factory +from revng.pypeline.storage.storage_provider import StorageProvider +from revng.pypeline.task.task import TaskArgument, TaskArgumentAccess +from revng.pypeline.utils.registry import get_registry, get_singleton + +logger = logging.getLogger(__name__) + + +def default_storage_provider(model_path: str) -> StorageProvider: + """ + While we figure out how to configure the storage provider, we use a default + SQLite3 storage provider. This is a temporary solution, and we should + eventually allow the user to configure the storage provider using an + environment variable or a configuration file. + """ + return storage_provider_factory(model_path, os.environ.get("PYPELINE_STORAGE")) + + +class RegistryChoice(click.Choice): + """A click.Choice that uses the registry of a given type, and + returns the actual object from the registry instead of just the string.""" + + def __init__( + self, ty: type, case_sensitive: bool = False, subclass_filter: Optional[type] = None + ) -> None: + self.ty = ty + self.registry: dict[str, type] = { # type: ignore[var-annotated] + k: v + for k, v in get_registry(ty).items() + if subclass_filter is None or issubclass(v, subclass_filter) + } + super().__init__( + choices=sorted(self.registry.keys()), + case_sensitive=case_sensitive, + ) + + def convert( + self, value: Any, param: Optional[click.Parameter], ctx: Optional[click.Context] + ) -> Any: + res = super().convert(value, param, ctx) + # Compared to a normal click.Choice, we need to return the actual object + # from the registry, not just the string + if isinstance(res, str): + if res not in self.registry: + raise ValueError("This should already be checked by click. ") + return self.registry[res] + return super().convert(value, param, ctx) + + +class EagerParsedPath(click.Path): + """ + A click.Path that does eager parsing, meaning that it will call your function during parsing. + This is useful for arguments that need to be parsed in order to provide + useful auto-completion or validation. + The cli and other arguments can retrieve the parsed value from the context object + using the name of the argument, like `ctx.obj['pipebox']`. + """ + + def __init__( + self, + name: str, + parser: Callable[[str], Any], + *args, + **kwargs, + ): + # Sensible defaults for our use case + kwargs.setdefault("exists", True) + kwargs.setdefault("dir_okay", False) + kwargs.setdefault("resolve_path", True) + super().__init__(*args, **kwargs) + self.parser = parser + self.name = name + + def convert( + self, value: Any, param: Optional[click.Parameter], ctx: Optional[click.Context] + ) -> Any: + res = super().convert(value, param, ctx) + if isinstance(res, str): + # If the value is a string, we parse it using the provided parser + res = self.parser(res) + if ctx is not None: + if ctx.obj is None: + ctx.obj = {} + if self.name in ctx.obj: + raise ValueError( + f"Argument `{self.name}` already set in context, " + "this is likely a bug in the code." + ) + # Store the parsed value in the context object + ctx.obj[self.name] = res + return res + + +class LazyGroup(click.Group): + """A click.Group that lazily loads subcommands from modules. + This allows for a more modular command line interface where subcommands + can be defined in separate modules and loaded only when needed. + This is good for performance, but it's especially needed for us to load the + pipebox before loading the subcommands, so the subcommands can assume that the + registries are already populated with the objects defined in the pipebox. + + This implementation is a slightly modified version of the one from + the click documentation: https://click.palletsprojects.com/en/stable/complex/ + """ + + # lazy_subcommands is a map of the form: + # + # {command-name} -> {module-name}:{command-object-name} + # + def __init__(self, *args, lazy_subcommands=None, **kwargs): + super().__init__(*args, **kwargs) + self.lazy_subcommands = lazy_subcommands or {} + + def list_commands(self, ctx): + base = super().list_commands(ctx) + lazy = sorted(self.lazy_subcommands.keys()) + return base + lazy + + def get_command(self, ctx, cmd_name): + if cmd_name in self.lazy_subcommands: + return self._lazy_load(cmd_name) + return super().get_command(ctx, cmd_name) + + def _lazy_load(self, cmd_name): + # Lazily loading a command, first get the module name and attribute name + import_path = self.lazy_subcommands[cmd_name] + modname, cmd_object_name = import_path.rsplit(":", 1) + # Do the import + mod = importlib.import_module(modname) + # Get the Command object from that module + cmd_object = getattr(mod, cmd_object_name) + # Check the result to make debugging easier + if not isinstance(cmd_object, click.Command): + raise ValueError( + f"Lazy loading of {import_path} failed by returning " "a non-command object" + ) + return cmd_object + + +def normalize_whitespace(text: str) -> str: + """ + Normalize whitespace in a string by removing leading and trailing + whitespace and replacing multiple spaces with a single space. + """ + text = re.sub(r"\s+", " ", text) + return text.strip() + + +def normalize_flag(name: str) -> str: + """ + Normalize a flag name by replacing spaces and underscores with + hyphens and converting it to lowercase. + """ + return normalize_whitespace(name).replace(" ", "-").replace("_", "-").lower() + + +def normalize_pos_arg_name(name: str) -> str: + """ + Normalize a positional argument name by replacing spaces and underscores + with hyphens and converting it to lowercase. + This is used for positional arguments that are not flags. + """ + return normalize_whitespace(name).replace(" ", "_").replace("-", "_").upper() + + +def build_arg_objects( + container_decl: ContainerDeclaration, +) -> Callable: + """ + A decorator that adds an argument to a click command for + the objects that the user wants in a specific container. + """ + arg_name = normalize_flag(container_decl.name) + object_id = container_decl.container_type.kind + return click.option( + f"--{arg_name}-objects", + metavar=f"/{object_id.__name__}1,/{object_id.__name__}2,...", + type=str, + help=( + f"The objects to require from container {arg_name.upper()}" + " as a comma-separated list of IDs. If not passed, all " + "objects will be requested." + ), + ) + + +def build_help_text( + args: Sequence[TaskArgument], + prologue: str = "", + epilogue: str = "", +) -> str: + """ + Build a standardized help text for a command. + """ + help_text = prologue + help_text += "\n\n\b\nArguments:" + help_text += "\n - [R] MODEL : Path - The path to the model file to use." + + for arg in args: + help_text += _build_help_line(arg) + + help_text += epilogue + return help_text + + +def _build_help_line(arg: TaskArgument) -> str: + if arg.access != TaskArgumentAccess.READ_WRITE: + access = "R" if arg.access == TaskArgumentAccess.READ else "W" + arg_name = normalize_pos_arg_name(arg.name) + line = f"\n - [{access}] {arg_name} : " + line += f"{arg.container_type.__name__} - " + line += normalize_whitespace(arg.help_text) + return line.rstrip() + # If the argument is read-write, we need to add both input and output + return _build_help_line( + TaskArgument( + name=f"{arg.name}-input", + container_type=arg.container_type, + help_text=arg.help_text, + access=TaskArgumentAccess.READ, + ) + ) + _build_help_line( + TaskArgument( + name=f"{arg.name}-output", + container_type=arg.container_type, + help_text="Like above, but for output.", + access=TaskArgumentAccess.WRITE, + ) + ) + + +def list_objects_for_container( + model: ReadOnlyModel, + arg_name: str, + kind: Kind, +): + """ + Print all available objects for a given container kind in the model. + """ + print(f"Available objects for `{arg_name}` kind: {kind.__name__}") + for obj in model.all_objects(kind): + print(f" - {obj}") + + +def compute_objects( + model: ReadOnlyModel, + arg_name: str, + kind: Kind, + kwargs: dict[str, str], +) -> ObjectSet: + """ + Check if the user provided a list of objects for the given + argument name, and if so, return an ObjectSet with those objects + deserialized. + Otherwise, return all objects of the given kind from the model. + """ + arg_name = normalize_flag(arg_name) + obj_id_ty = get_singleton(ObjectID) # type: ignore[type-abstract] + if f"{arg_name}_objects" in kwargs: + objects = kwargs.get(f"{arg_name}_objects", "") + if objects: + return ObjectSet( + kind=kind, + objects={obj_id_ty.deserialize(obj) for obj in objects.split(",") if obj.strip()}, + ) + return model.all_objects(kind) diff --git a/python/revng/pypeline/container.py b/python/revng/pypeline/container.py new file mode 100644 index 000000000..b546032cc --- /dev/null +++ b/python/revng/pypeline/container.py @@ -0,0 +1,195 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +import json +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Annotated, Dict, Generator, Optional, Tuple, Type + +from .object import Kind, ObjectID, ObjectSet +from .utils.cabc import ABC, abstractmethod +from .utils.registry import get_singleton + +ContainerID = Annotated[ + str, + """ + This is used to identify a container in a savepoint. + This is the name of the container. + """, +] + +ConfigurationId = Annotated[ + str, + """ + From a pipeline node, this represent the static and runtime configuration of + this node, and all its dependencies. This is used to index the storage provider + in SavePoints. + """, +] + +Configuration = Annotated[ + str, + "The runtime-configuration of a pipe, most commonly it will be a json / yaml string that" + "the pipe will parse and use to configure itself." + "THIS IS DIFFERENT FROM ConfigurationId.", +] + + +@dataclass(slots=True, frozen=True) +class ContainerDeclaration: + """ + A ContainerDeclaration represents a container when describing a pipeline. + It has a name and a type. + + Not to be confused with a Container, which is an instance of the container_type. + """ + + name: str + container_type: Type[Container] + + def instance(self) -> Container: + return self.container_type() + + +InvalidationList = list[Tuple[ConfigurationId, ContainerDeclaration, ObjectID]] + + +def group_by_container( + invalidation_list: InvalidationList, +) -> Generator[Tuple[ConfigurationId, ContainerDeclaration, ObjectSet], None, None]: + if len(invalidation_list) == 0: + return + + first_configuration_id, first_container, _ = invalidation_list[0] + last = (first_configuration_id, first_container) + + def new_object_list(): + return ObjectSet(first_container.container_type.kind, set()) + + object_list: ObjectSet = new_object_list() + for configuration_id, container, obj in invalidation_list: + if last != (configuration_id, container): + # We changed container. Yield what we accumulated so far and prepare a new list + yield (last[0], last[1], object_list) + last = (configuration_id, container) + object_list = new_object_list() + + # Record the current object + object_list.add(obj) + + # Yield the last group + yield (last[0], last[1], object_list) + + +class Container(ABC): + """ + A Container contains objects of a certain kind. + """ + + kind: Kind + + def __init__(self): + """ + This constructor just makes it explicit that a container should be + able to be initialized without any arguments. + """ + + @abstractmethod + def objects(self) -> ObjectSet: + pass + + @abstractmethod + def deserialize(self, data: Mapping[ObjectID, bytes]) -> None: + """ + Ingest data from a serialized format into this container. + This is used to **add** cached objects to this container. + """ + + @abstractmethod + def serialize(self, objects: Optional[ObjectSet] = None) -> Mapping[ObjectID, bytes]: + """ + Dump objects from this container into a serialized format. + If objects is provided, only those objects will be dumped. + If not, all objects in the container will be dumped. + """ + + @classmethod + @abstractmethod + def mime_type(cls) -> str: + """ + The mime type of the serialized format of this container. + This is used to inform the storage provider about the type of data + it will be storing. + """ + + @classmethod + def is_text(cls) -> bool: + """ + Returns if the serialized format of this container is just + text (e.g. JSON, YAML, etc.) or binary. + This is used to improve transmission performance by avoiding + unnecessary encoding/decoding steps. + """ + return cls.mime_type().startswith("text/") or cls.mime_type() in { + "application/json", + "application/xml", + "application/x-yaml", + } + + @abstractmethod + def verify(self) -> bool: + pass + + def contains_all(self, obj: ObjectSet) -> bool: + """ + Check if an object set is fully contained in this container. + This is a default implementation, the container can probably do it much + more efficiently, so you are supposed to override this method. + """ + assert ( + self.kind == obj.kind + ), f"Container {self} has kind {self.kind}, but the object set has kind {obj.kind}." + return self.objects().issubset(obj) + + +def load_container( + container_type: type[Container], + path: str, +) -> Container: + """ + Load a container from a serialized format. + This is used to **load** cached objects into this container. + """ + container = container_type() + obj_id_ty = get_singleton(ObjectID) # type: ignore[type-abstract] + with open(path, "r", encoding="utf-8") as f: + container.deserialize( + {obj_id_ty.deserialize(k): bytes.fromhex(v) for k, v in json.load(f).items()} + ) + return container + + +def dump_container( + container: Container, + path: str, +) -> None: + """ + Dump a container into a serialized format. + This is used to **save** cached objects from this container. + """ + with open(path, "w", encoding="utf-8") as f: + json.dump( + {k.serialize(): v.hex() for k, v in container.serialize().items()}, + f, + indent=4, + sort_keys=True, + ) + + +ContainerSet = Annotated[ + Dict[ContainerDeclaration, Container], + """A set of bindings between container declarations and container instances.""", +] diff --git a/python/revng/pypeline/graph.py b/python/revng/pypeline/graph.py new file mode 100644 index 000000000..32ee4787b --- /dev/null +++ b/python/revng/pypeline/graph.py @@ -0,0 +1,97 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from dataclasses import dataclass, field +from typing import List, Set, TypeAlias + + +def escape(string: str) -> str: + return string.replace("\\", "\\\\").replace('"', '\\"') + + +@dataclass(slots=True) +class Node: + label: str + completed: bool = False + entries: List[str] = field(default_factory=list) + + def name(self) -> str: + return f"node_{id(self)}" + + def to_graphviz(self) -> str: + result = "" + result += f""" {escape(self.name())} [""" + result += "label=<\n" + result += """ \n""" + + cell_properties = "" + if self.completed: + cell_properties = ' bgcolor="lightgreen"' + + result += f""" {escape(self.label)}\n""" + + for index, entry in enumerate(self.entries): + result += f""" \n""" + result += "
{escape(entry)}
\n" + result += " >];\n" + result += "\n" + return result + + def __hash__(self) -> int: + return hash(self.name()) + + +@dataclass(slots=True) +class Edge: + source: Node + destination: Node + source_port: int = 0 + destination_port: int = 0 + head_label: str = "" + tail_label: str = "" + + def to_graphviz(self) -> str: + result = "" + result += " " + result += f"{escape(self.source.name())}:entry_{self.source_port}" + result += " -> " + result += f"{escape(self.destination.name())}:entry_{self.destination_port}" + if self.head_label or self.tail_label: + result += " [" + if self.head_label: + result += f'headlabel="{escape(self.head_label)}"' + if self.tail_label: + if self.head_label: + result += "," + result += f'taillabel="{escape(self.tail_label)}"' + result += "]" + result += ";\n" + return result + + +class Graph: + """A graph data structure for rendering purposes.""" + + Node: TypeAlias = Node + Edge: TypeAlias = Edge + + def __init__(self) -> None: + self.nodes: Set[Node] = set() + self.edges: Set[Edge] = set() + + def is_tree(self) -> bool: + # TODO: implement + return True + + def to_graphviz(self) -> str: + result = "digraph structs {\n" + result += "rankdir = LR;\n" + result += "node [shape=plaintext];\n" + result += "\n\n" + for node in self.nodes: + result += node.to_graphviz() + for edge in self.edges: + result += edge.to_graphviz() + result += "\n}\n" + return result diff --git a/python/revng/pypeline/main.py b/python/revng/pypeline/main.py new file mode 100644 index 000000000..f917acd21 --- /dev/null +++ b/python/revng/pypeline/main.py @@ -0,0 +1,102 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +import importlib.util +import logging +import os +import sys +from pathlib import Path +from typing import Sequence + +import click + +from . import initialize_pypeline +from .cli.utils import EagerParsedPath, LazyGroup + +logger = logging.getLogger("pypeline") +logger.setLevel(logging.DEBUG) +logger.addHandler(logging.StreamHandler(sys.stderr)) + + +def import_pipebox(module_path: str, is_complete: bool) -> object: + """Import a module from a path. This is used to import the pipebox file. + Args: + env (dict[str, str]): The environment variables. + module_path (str): The path to the module. + is_complete (bool): If True, raise an error if the module is not found. + """ + # Absolute path to the module + module_abspath = Path(module_path).resolve() + if not module_abspath.exists(): + # This is a small trick to allow generating the module auto-complete + # without having a pipebox file + if is_complete: + return object() + logger.error( + ( + "Pipebox file `%s` does not exist. Either set it using the " + "PIPEBOX env var, or pass the --pipebox option." + ), + module_abspath, + ) + sys.exit(1) + # We guess that the module name is the file name without the extension + module_name: str = module_abspath.stem + # Dynamic import of the pipebox module + spec = importlib.util.spec_from_file_location(module_name, str(module_abspath)) + if spec is None: + if is_complete: + return object() + logger.error("Could not load module `%s` from `%s`", module_name, module_abspath) + sys.exit(1) + module = importlib.util.module_from_spec(spec) + if spec.loader is None: + if is_complete: + return object() + logger.error("Could not load module `%s` from `%s`", module_name, module_abspath) + sys.exit(1) + # Execute the module to load it + spec.loader.exec_module(module) + # Initialize the pypeline as we just imported the pipebox + initialize_pypeline() + return module + + +@click.group( + cls=LazyGroup, + lazy_subcommands={ + "pipeline": "revng.pypeline.cli.pipeline:pipeline", + "project": "revng.pypeline.cli.project:project", + }, +) +@click.option( + "--pipebox", + type=EagerParsedPath( + name="pipebox", + parser=lambda path: import_pipebox(path, "_PYPE_COMPLETE" in os.environ), + ), + help=( + "Path to the pipebox file. Defaults to the `PIPEBOX` environment " + "variable, then 'pipebox.py'." + ), + default=os.environ.get("PIPEBOX", "pipebox.py"), + expose_value=False, +) +def cli(): + pass + + +def main(args: Sequence[str]) -> None: + # pylint: disable=E1120 no-value-for-parameter + # This is ok as click will pass the pipebox argument automatically + cli(args=args) + + +def run(): + """Run the pipeline from the command line using the shell environment.""" + main(sys.argv[1:]) + + +if __name__ == "__main__": + run() diff --git a/python/revng/pypeline/model.py b/python/revng/pypeline/model.py new file mode 100644 index 000000000..419ff9084 --- /dev/null +++ b/python/revng/pypeline/model.py @@ -0,0 +1,149 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +from typing import Annotated, Self, Set + +from .object import Kind, ObjectID, ObjectSet +from .utils.cabc import ABC, abstractmethod +from .utils.registry import get_singleton + +ModelPath = Annotated[str, "A string that represents a path in the model tree."] +ModelPathSet = Set[ModelPath] + + +class Model(ABC): + """ + Model is an (abstract) class representing a document that configures the what a certain pipeline + should produce. + The document is a tree composed by dictionaries, list and scalars. + It can be navigated via ModelPath. + + From the model, it's possible to enumerate all the objects of a certain Kind. + + The model can be changed, a change in the model might need to the invalidation of certain + objects that have been previously produced. + """ + + @abstractmethod + def diff(self, other: Self) -> ModelPathSet: + raise NotImplementedError() + + @abstractmethod + def clone(self) -> Self: + raise NotImplementedError() + + @abstractmethod + def children(self, obj: ObjectID, kind: Kind) -> ObjectSet: + raise NotImplementedError() + + @abstractmethod + def __eq__(self, other: object) -> bool: + raise NotImplementedError() + + def all_objects(self, kind: Kind) -> ObjectSet: + """ + Returns all the objects of a certain kind in the model. + """ + obj_id_ty: type[ObjectID] = get_singleton(ObjectID) # type: ignore[type-abstract] + root: ObjectID = obj_id_ty.root() + return ObjectSet(root.kind(), self.children(root, kind).objects) + + def move_to_kind(self, objects: ObjectSet, destination_kind: Kind) -> ObjectSet: + if not objects: + # If the objects set is empty, we can return an empty set of the destination kind + return ObjectSet(destination_kind) + + source_kind = objects.kind + + match destination_kind.relation(source_kind): + case (Kind.Relation.SAME, _): + return objects.clone() + case (Kind.Relation.UNRELATED, _): + return ObjectSet(destination_kind) + case (Kind.Relation.ANCESTOR, path): + assert isinstance(path, list) + result = ObjectSet(destination_kind) + for obj in objects.objects: + for _ in range(len(path) - 1): + parent = obj.parent() + assert parent is not None, ( + f"Object {obj} of kind {source_kind} has no parent, " + f"but we are trying to move it to {destination_kind} which " + "is an ancestor." + ) + obj = parent + result.add(obj) + return result + case (Kind.Relation.DESCENDANT, path): + assert isinstance(path, list) + + result = objects + + # For each child kind + for child_kind in path[:-1]: + # Create a new object set and populate it with the children from the previous + # iteration + children = ObjectSet(kind=child_kind) + + for obj in result: + children.update(self.children(obj, child_kind)) + + # Make children become the source for the next iteration + result = children + + assert result.kind is destination_kind + return result + + assert False, f"Unhandled relation between kinds: {source_kind} -> {destination_kind}" + + @abstractmethod + def serialize(self) -> bytes: + pass + + @abstractmethod + def deserialize(self, data: bytes): + pass + + @classmethod + def is_text(cls) -> bool: + """ + Returns True if the serialized model is a text file, False otherwise. + This is used to determine if the model should be treated as a text file or not. + """ + return False + + +class ReadOnlyModel[M: Model]: + """ + A wrapper around the Model ensuring no changes can be made. + """ + + def __init__(self, inner: M): + self._context: M = inner + + def diff(self, other: ReadOnlyModel[M]) -> ModelPathSet: + return self._context.diff(other._context) # pylint: disable=protected-access + + def clone(self) -> M: + return self._context.clone() + + def children(self, obj: ObjectID, kind: Kind) -> ObjectSet: + return self._context.children(obj, kind) + + def all_objects(self, kind: Kind) -> ObjectSet: + return self._context.all_objects(kind) + + def move_to_kind(self, objects: ObjectSet, destination_kind: Kind) -> ObjectSet: + return self._context.move_to_kind(objects, destination_kind) + + def downcast(self) -> M: + """ + Returns the inner model, allowing to use it as a regular model. + The user should be careful not to modify the model. + While we could convert this to a context manager, real users will be + written in C++ where we can enforce the immutability of the model. + """ + return self._context diff --git a/python/revng/pypeline/object.py b/python/revng/pypeline/object.py new file mode 100644 index 000000000..b36147dd9 --- /dev/null +++ b/python/revng/pypeline/object.py @@ -0,0 +1,313 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +from collections.abc import MutableSet +from dataclasses import dataclass, field +from enum import Enum + +# builtin since python 3.9 +from graphlib import TopologicalSorter +from typing import Iterator, Optional, Sequence, Set, cast + +from .graph import Graph +from .utils.cabc import ABC, abstractmethod + + +class Kind(ABC): + __name__: str + # Class attributes set in __init_subclass__ + _ranks: dict[Kind, int] + _children: dict[Kind, set[Kind]] + _root: Kind | None + + # These two methods define the hierarchy + @classmethod + @abstractmethod + def kinds(cls) -> list[Kind]: + """Return a list of all the kinds""" + raise NotImplementedError() + + @abstractmethod + def parent(self) -> Kind | None: + """Get the parent of this kind, or None if root""" + raise NotImplementedError() + + @abstractmethod + def __hash__(self) -> int: + """Standard python hash, needed to put Kinds in sets and dicts""" + raise NotImplementedError() + + # Back and forth from str to kind for storage provider and CLI + @abstractmethod + def serialize(self) -> str: + """Convert the kind to a string representation that can be later be deserialized""" + raise NotImplementedError() + + @classmethod + @abstractmethod + def deserialize(cls, value: str) -> Kind: + """Convert a string representation back to a kind""" + raise NotImplementedError() + + # Methods we can provide concrete implementations based on the above methods + def __eq__(self, other) -> bool: + return hash(self) == hash(other) + + def __str__(self) -> str: + return self.serialize() + + def __repr__(self) -> str: + return self.serialize() + + @classmethod + def _init_type(cls): + """ + Verify that `all` and `parent` are at reasonable and compute useful + things like ranks, children, and the root. + """ + kinds = set(cls.kinds()) + assert len(kinds) > 0, "There must be at least one kind" + cls._ranks: dict[Kind, int] = {} + """Cache of the rank of each kind""" + cls._children: dict[Kind, set[Kind]] = {} + """Cache of the children of each kind""" + cls._root: Kind | None = None + """The root of the kinds""" + for kind in kinds: + # Check that it's reasonable + assert isinstance(kind, Kind) + parent = kind.parent() + if parent is None: + assert cls._root is None, f"Found two roots! {cls._root} and {kind}" + cls._root = kind + cls._ranks[cls._root] = 0 + else: + cls._children.setdefault(parent, set()).add(kind) + assert parent in kinds, ( + "The parent of each kind should be a known kind. " + f"Got parent {parent} which is not in the " + f"known kinds '{kinds}'" + ) + assert ( + cls._root is not None + ), "Could not find a root, there must be a loop in the hierarchy." + # Assign ranks + ts = TopologicalSorter(cls._children) + for kind in reversed(list(ts.static_order())): + parent = kind.parent() + if parent is None: + continue + cls._ranks[kind] = cls._ranks[parent] + 1 + + def children(self) -> set[Kind]: + """ + Return all the children of this kind, this is a generic impl + and the implementer can probably write a more efficient one. + """ + return self.__class__._children.get(self, set()) + + def rank(self) -> int: + """ + Return the distance of the current kind from the root. + """ + return self.__class__._ranks[self] + + @classmethod + def root(cls) -> Kind: + # We already check in the __init_subclass__ that there is a root + return cast(Kind, cls._root) + + def is_subkind_of(self, other: Kind) -> bool: + """ + Returns whether the current kind is equal + to `other` or it's a descendent of `other` + """ + if self == other: + return True + parent = self.parent() + if parent is None: + return False + return parent.is_subkind_of(other) + + @classmethod + def graph( + cls, + ) -> Graph: + """ + Returns a graph of the full kinds hierarchy for printing porpouses + """ + graph = Graph() + nodes = {} + # Create all nodes + for kind in cls.kinds(): + node = Graph.Node(kind.__name__) + graph.nodes.add(node) + nodes[kind] = node + # Create all edges + for kind in cls.kinds(): + parent = kind.parent() + # The root can't have arcs + if parent is None: + continue + graph.edges.add(Graph.Edge(nodes[parent], nodes[kind])) + return graph + + class Relation(Enum): + SAME = 0 + ANCESTOR = 1 + DESCENDANT = 2 + UNRELATED = 3 + + def relation(self, other: Kind) -> tuple[Relation, list[Kind] | None]: + """ + Returns the relation between the two kinds, and, if related, returns + the path to get from the ancestor to the descendent. + """ + # Easy case + if self == other: + return Kind.Relation.SAME, None + # Same rank but different means always unrelated + if self.rank() == other.rank(): + return Kind.Relation.UNRELATED, None + # Start from descendant and raise up to ancestor + ancestor = min(self, other, key=lambda x: x.rank()) + descendant = max(self, other, key=lambda x: x.rank()) + path = [descendant] + while descendant != ancestor: + parent = descendant.parent() + # If we got to the root, the other node must be on a different + # branch, thus unrelated + if parent is None: + return Kind.Relation.UNRELATED, None + path.append(parent) + descendant = parent + + if self == ancestor: + return Kind.Relation.ANCESTOR, path[::-1] + else: + return Kind.Relation.DESCENDANT, path + + +class ObjectID(ABC): + # Needed for model's `move_to_kind` + @abstractmethod + def kind(self) -> Kind: + """Return the kind of this object""" + raise NotImplementedError() + + @abstractmethod + def parent(self) -> Optional[ObjectID]: + """Return the parent object of this object, or None if root""" + raise NotImplementedError() + + @classmethod + @abstractmethod + def root(cls) -> ObjectID: + """Return an instance of the root object of this hierarchy""" + raise NotImplementedError() + + @abstractmethod + def __hash__(self) -> int: + """Standard python hash, needed to put ObjectIDs in sets and dicts""" + raise NotImplementedError() + + @abstractmethod + def serialize(self) -> str: + """Convert the object to a string representation that can be later be deserialized""" + raise NotImplementedError() + + @classmethod + @abstractmethod + def deserialize(cls, obj: str) -> ObjectID: + """Convert a string representation back to an object""" + raise NotImplementedError() + + def __eq__(self, other) -> bool: + return hash(self) == hash(other) + + def __str__(self) -> str: + return self.serialize() + + def __repr__(self) -> str: + return repr(str(self)) + + +@dataclass(slots=True) +class ObjectSet(MutableSet[ObjectID]): + """ + A list of objects with the same kind. + """ + + kind: Kind + """ + The kind of all the objects in this set. + """ + + objects: Set[ObjectID] = field(default_factory=set) + """ + The objects in this set. + TODO!: the user shouldn't be able to access this directly, but rather + through the methods of this class. + """ + + @staticmethod + def from_list(seq: Sequence[ObjectID]): + assert len(seq) > 0 + result = ObjectSet(seq[0].kind(), set(seq)) + return result + + def clone(self) -> ObjectSet: + result = ObjectSet(self.kind) + result.objects = set(self.objects) + return result + + def __post_init__(self): + for obj in self.objects: + assert isinstance(obj, ObjectID), f"Expected ObjectID, got {obj}" + assert obj.kind() == self.kind + + def __contains__(self, obj: object) -> bool: + assert isinstance(obj, ObjectID) + assert obj.kind() == self.kind + return obj in self.objects + + def __iter__(self) -> Iterator[ObjectID]: + return self.objects.__iter__() + + def __len__(self) -> int: + return len(self.objects) + + def __eq__(self, other: object) -> bool: + if not isinstance(other, ObjectSet): + return False + if self.kind != other.kind: + return False + return self.objects == other.objects + + def __repr__(self) -> str: + return f"ObjectSet(kind={self.kind.serialize()}, objects={self.objects})" + + def add(self, value: ObjectID): + assert value.kind() == self.kind + self.objects.add(value) + + def discard(self, value: ObjectID): + assert value.kind() == self.kind + self.objects.discard(value) + + def update(self, *others: ObjectSet): + for other in others: + if not isinstance(other, ObjectSet): + raise TypeError(f"Expected ObjectSet, got {type(other)}") + if other.kind != self.kind: + raise ValueError( + f"Cannot update ObjectSet of kind {self.kind} with ObjectSet" + f" of kind {other.kind}." + ) + self.objects.update(other.objects) + + def issubset(self, other: ObjectSet) -> bool: + return self <= other diff --git a/python/revng/pypeline/pipeline.py b/python/revng/pypeline/pipeline.py new file mode 100644 index 000000000..8bae47da4 --- /dev/null +++ b/python/revng/pypeline/pipeline.py @@ -0,0 +1,415 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +import logging +from dataclasses import dataclass +from itertools import chain +from typing import Dict, Generator, Generic, List, Mapping, Optional, Set, TypeVar + +import yaml + +from .analysis import AnalysisBinding +from .container import Container, ContainerDeclaration +from .graph import Graph +from .model import Model, ReadOnlyModel +from .object import ObjectID, ObjectSet +from .pipeline_node import PipelineConfiguration, PipelineNode +from .schedule.schedule import Schedule +from .schedule.scheduled_task import ScheduledTask +from .storage.storage_provider import SavePointsRange, StorageProvider +from .task.pipe import Pipe +from .task.requests import Requests +from .task.savepoint import SavePoint +from .task.task import TaskArgumentAccess +from .utils.default_dict_from_key import DefaultDictFromKey +from .utils.registry import get_singleton + +logger = logging.getLogger(__name__) + + +@dataclass(frozen=True, slots=True) +class Artifact: + """ + An artifact is a container in a certain point of the pipeline with some extra + metadata, such as a name. + + It's designed to mark interesting results in the pipeline, so that users an + easily obtain them. + """ + + name: str + node: PipelineNode + container: ContainerDeclaration + + +C = TypeVar("C", bound=Model) + + +class Pipeline(Generic[C]): + """ + A pipeline is a tree of tasks. + + Given a set of requests, a model and a configuration of the pipes, it can produce a schedule + that fulfills the requests. + """ + + __slots__ = ("declarations", "root", "artifacts", "analyses") + + def __init__( + self, + declarations: Set[ContainerDeclaration], + root: PipelineNode, + artifacts: Optional[set[Artifact]] = None, + analyses: Optional[set[AnalysisBinding]] = None, + ): + self.root = root + self.declarations = set(declarations) + + self.artifacts: Mapping[str, Artifact] = {} + """ + The artifacts, indexed by their name for easy access. + """ + for artifact in artifacts or set(): + if artifact.name in self.artifacts: + raise ValueError( + f"Artifact {artifact.name} is defined multiple times in the pipeline" + ) + self.artifacts[artifact.name] = artifact + + self.analyses: Mapping[str, AnalysisBinding] = {} + """ + The analyses, indexed by their name for easy access. + """ + for analysis in analyses or set(): + if analysis.analysis.name in self.analyses: + raise ValueError( + f"Analysis {analysis.analysis.name} is defined multiple times in the pipeline" + ) + self.analyses[analysis.analysis.name] = analysis + + Pipeline.assign_savepoint_ranges(root) + + id_index = 0 + # Set the dependencies field + for node in self.walk_pipeline(stable=True): + node.id = id_index + id_index += 1 + + if node is self.root: + node.pipe_dependencies = set() + continue + + node.pipe_dependencies = set( + chain.from_iterable(n.pipe_dependencies for n in node.predecessors) + ) + if isinstance(node.task, Pipe): + node.pipe_dependencies.add(node.task) + + for name, artifact in self.artifacts.items(): + if name != artifact.name: + raise ValueError( + f"Artifact name {artifact.name} does not match the key " + f"{name} in the artifacts map." + ) + + @staticmethod + def assign_savepoint_ranges(node: PipelineNode, current_id: int = 0) -> int: + """ + Assigns savepoint ranges to the nodes in the pipeline tree. + A savepoint range is a pair of integers that represent the start and end of + the savepoint in a subtree. The idea is to deduplicate things common to a + savepoint and all its children, so that we can efficiently represent a + subtree of savepoints as a continuous range of integers. + """ + # We ID only the savepoints, so we increment the ID only if the current node is a savepoint. + # This implies that, while the ids start at 0, the first savepoint will have id 1. + # This is needed in the case the first node is a pipe, which is not a savepoint + if isinstance(node.task, SavePoint): + current_id += 1 + # Save the id on preorder visit + start = current_id + # Recurse on the children, but do it in a deterministic way + end = current_id + for child in node.sorted_successors(): + current_id = Pipeline.assign_savepoint_ranges(child, current_id) + assert child.savepoint_range is not None, ( + f"Child {child.task.name} does not have a savepoint range assigned:" + f" {child.savepoint_range}" + ) + end = max(end, child.savepoint_range.end) + # The end of the savepoint range is assigned on postorder visit. + # Therefore the end is inclusive + assert node.savepoint_range is None, ( + f"SavePoint {node.task.name} already has a savepoint range assigned" + ": {node.savepoint_range}" + ) + node.savepoint_range = SavePointsRange(start, end) + # Return the new id + return current_id + + def walk_pipeline( + self, start: Optional[PipelineNode] = None, forward: bool = True, stable: bool = False + ) -> Generator[PipelineNode, None, None]: + """BFS walk of pipeline nodes""" + to_visit: List[PipelineNode] = [start or self.root] + visited: Set[PipelineNode] = set() + + if forward: + + def successors(node): + if not stable: + return node.successors + else: + return node.sorted_successors() + + else: + + def successors(node): + assert not stable + return node.predecessors + + while len(to_visit) > 0: + node = to_visit.pop() + yield node + visited.add(node) + for child_node in successors(node): + if child_node not in visited: + to_visit.append(child_node) + + def graph(self) -> Graph: + """A graph for debugging purposes.""" + + graph = Graph() + + nodes_map: Dict[PipelineNode, Graph.Node] = {} + + def get_node(node: PipelineNode) -> Graph.Node: + if node not in nodes_map: + new_node = Graph.Node(node.task.name) + for argument in node.arguments: + new_node.entries.append(argument.name) + + nodes_map[node] = new_node + graph.nodes.add(new_node) + + return nodes_map[node] + + for node in self.walk_pipeline(): + graph_node = get_node(node) + node_inputs: List[ContainerDeclaration] = list(node.arguments) + + for predecessor in node.predecessors: + for source_index, argument in enumerate(predecessor.task.arguments): + if argument.access == TaskArgumentAccess.READ or argument not in node_inputs: + continue + + destination_index = node_inputs.index(argument) + graph.edges.add( + Graph.Edge( + get_node(predecessor), + graph_node, + source_port=source_index, + destination_port=destination_index, + ) + ) + + return graph + + def schedule( + self, + model: ReadOnlyModel, + target_node: PipelineNode, + requests: Requests, + pipeline_configuration: PipelineConfiguration, + storage_provider: StorageProvider, + ) -> Schedule: + tasks: DefaultDictFromKey[PipelineNode, ScheduledTask] = DefaultDictFromKey(ScheduledTask) + # The pipeline is a tree, so we can just unroll the predecessors, + # When we parallelize, we will make a subclass that overrides this method, + # and probably it will first call it to produce the initial schedule and + # then add the parallelization logic + + node: PipelineNode = target_node + node_outgoing_requests = requests + + while not node_outgoing_requests.empty(): + logger.debug("Scheduling node %s", node) + logger.debug("Outgoing requests: %s", node_outgoing_requests) + orig = repr(node_outgoing_requests) + # Each node should remove the requests it can handle + # and add the requests it needs to satisfy the task + node_ingoing_requests = node.prerequisites_for( + model=model, + requests=node_outgoing_requests, + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + assert orig == repr(node_outgoing_requests), ( + f"Node {node} modified the outgoing requests, which is not allowed. " + f"Original: {orig}, modified: {node_outgoing_requests}" + ) + logger.debug("Computed Ingoing requests: %s", node_ingoing_requests) + # Store the computed requests so that we can use them in the run method + tasks[node].request(node_ingoing_requests, node_outgoing_requests) + # If the node has no predecessors, we are done, but + # we need to check that it has no requests left + if not node.predecessors: + assert node_ingoing_requests.empty(), ( + f"Node {node} has no predecessors, but it still has " + f"requests: {node_ingoing_requests}" + ) + break + + assert len(node.predecessors) == 1, ( + f"Node {node} has multiple predecessors, but we assume a tree structure. " + f"Predecessors: {node.predecessors}" + ) + predecessor = node.predecessors[0] + tasks[node].depends_on.append(tasks[predecessor]) + # Recurse on THE predecessor, its outgoing requests will be the + # ingoing requests of the current node + node = predecessor + node_outgoing_requests = node_ingoing_requests + + return Schedule(self.declarations, tasks[target_node], pipeline_configuration) + + def get_artifact( + self, + model: ReadOnlyModel, + artifact: Artifact, + requests: ObjectSet, + pipeline_configuration: PipelineConfiguration, + storage_provider: StorageProvider, + ) -> Container: + schedule = self.schedule( + model=model, + target_node=artifact.node, + requests=Requests({artifact.container: requests}), + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + return schedule.run( + model=model, + storage_provider=storage_provider, + )[artifact.container] + + def run_analysis( + self, + model: ReadOnlyModel, + analysis_name: str, + requests: Requests, + analysis_configuration: str, + pipeline_configuration: PipelineConfiguration, + storage_provider: StorageProvider, + ) -> Model: + """ + Run an analysis on the pipeline, given a model and a set of requests. + The analysis will return the new potentially modified model, and set it + in the storage provider. + """ + if analysis_name not in self.analyses: + raise ValueError(f"Analysis {analysis_name} not found in the pipeline") + analysis_info = self.analyses[analysis_name] + + if len(requests) != len(analysis_info.bindings): + raise ValueError( + f"Expected {len(analysis_info.bindings)} requests for analysis " + f"{analysis_name}, but got {len(requests)}: {requests}" + ) + + for req in requests: + if req not in analysis_info.bindings: + raise ValueError( + f"Request {req} but it's not compatible with in the " + f"analysis bindings: {analysis_info.bindings}" + ) + + schedule = self.schedule( + model=model, + target_node=analysis_info.node, + requests=requests, + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + all_containers = schedule.run(model=model, storage_provider=storage_provider) + # The analysis modifies the model, but we need to invalidate + # the changes, so we make a clone of the model. + # This also allows to keep the original model intact + # in case the analysis fails + new_model = model.clone() + analysis_info.analysis.run( + model=new_model, + containers=[all_containers[decl] for decl in analysis_info.bindings], + incoming=[ + requests.get(decl, ObjectSet(decl.container_type.kind, set())) + for decl in analysis_info.bindings + ], + configuration=analysis_configuration, + ) + storage_provider.invalidate(ReadOnlyModel(new_model).diff(model)) + storage_provider.set_model(new_model.serialize()) + return new_model + + def deserialize_schedule(self, schedule: str) -> Schedule: + schedule_dict = yaml.safe_load(schedule) + declarations = set() + for container in schedule_dict["containers"]: + for declaration in self.declarations: + if container["name"] == declaration.name: + declarations.add(declaration) + break + else: + raise ValueError() + + pipeline_nodes = list(self.walk_pipeline(stable=True)) + container_map = {x.name: x for x in self.declarations} + + obj_id_ty = get_singleton(ObjectID) # type: ignore[type-abstract] + configuration = {} + scheduled_tasks: list[ScheduledTask] = [] + for task in schedule_dict["tasks"]: + pipeline_node: PipelineNode = pipeline_nodes[task["node_id"]] + outgoing = Requests() + incoming = Requests() + + if task["type"] == "Pipe": + assert isinstance(pipeline_node.task, Pipe) + assert task["name"] == pipeline_node.task.name + + configuration[pipeline_node.task] = task["dynamic_config"] + for index, arg in enumerate(task["args"]): + container_declaration = pipeline_node.bindings[index] + assert container_declaration.name == arg["name"] + container_kind = container_declaration.container_type.kind + + incoming[container_declaration] = ObjectSet( + container_kind, {obj_id_ty.deserialize(x) for x in arg["incoming"]} + ) + outgoing[container_declaration] = ObjectSet( + container_kind, {obj_id_ty.deserialize(x) for x in arg["outgoing"]} + ) + elif task["type"] == "SavePoint": + assert isinstance(pipeline_node.task, SavePoint) + assert task["name"] == pipeline_node.task.name + + for container in task["containers"]: + container_declaration = container_map[container["name"]] + container_kind = container_declaration.container_type.kind + + incoming[container_declaration] = ObjectSet( + container_kind, {obj_id_ty.deserialize(x) for x in container["incoming"]} + ) + outgoing[container_declaration] = ObjectSet( + container_kind, {obj_id_ty.deserialize(x) for x in container["outgoing"]} + ) + else: + raise ValueError(f"Unknown task type: {task['type']}") + + depends_on = [scheduled_tasks[i] for i in task["depends_on"]] + scheduled_task = ScheduledTask(pipeline_node, False, outgoing, incoming, depends_on) + scheduled_tasks.append(scheduled_task) + + return Schedule(declarations, scheduled_tasks[-1], configuration) diff --git a/python/revng/pypeline/pipeline_node.py b/python/revng/pypeline/pipeline_node.py new file mode 100644 index 000000000..4ab9430e7 --- /dev/null +++ b/python/revng/pypeline/pipeline_node.py @@ -0,0 +1,246 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +from hashlib import sha256 +from typing import TYPE_CHECKING, Annotated, List, Mapping, Optional, Sequence, Set, Union +from typing import overload + +from .container import Configuration, ConfigurationId, ContainerDeclaration, ContainerSet +from .model import ReadOnlyModel +from .object import ObjectSet +from .storage.storage_provider import SavePointsRange, StorageProvider +from .task.pipe import Pipe +from .task.requests import Requests +from .task.savepoint import SavePoint +from .task.task import ObjectDependencies, TaskArgumentAccess + +if TYPE_CHECKING: + from _typeshed import SupportsRichComparison + + +PipelineConfiguration = Annotated[ + Mapping[Pipe, Configuration], + """A configuration of a pipeline, specifying, for each Pipe, its runtime configuration""", +] + + +def deterministic_sort_key(node: PipelineNode) -> SupportsRichComparison: + """ + A deterministic sort key for a PipelineNode, used to ensure that the + savepoint ranges are assigned in a deterministic order. + """ + return ( + # Arguments of the task + tuple(node.arguments), + # Sorted list of outdegree of successors to try to embed the shape + # of the pipeline in the sort key + tuple(sorted(len(succ.successors) for succ in node.successors)), + node.task.name, + ) + + +class PipelineNode: + """ + A node of the tree of elements composing a pipeline. + + Mostly a wrapper around a Task. + """ + + __slots__ = ( + "successors", + "predecessors", + "task", + "bindings", + "reverse_bindings", + "pipe_dependencies", + "savepoint_range", + "id", + ) + + @overload + def __init__(self, task: Pipe, bindings: Sequence[ContainerDeclaration]): ... + + @overload + def __init__(self, task: SavePoint): ... + + def __init__( + self, + task: Union[Pipe, SavePoint], + bindings: Sequence[ContainerDeclaration] | None = None, + ): + """ + The bindings are used to map between the pipeline declarations and the + task arguments and are needed only for Pipe tasks as SavePoints only + live in the pipeline namespace. + """ + self.successors: List[PipelineNode] = [] + self.predecessors: List[PipelineNode] = [] + self.task = task + self.pipe_dependencies: Set[Pipe] = set() + self.bindings: list[ContainerDeclaration] = [] + self.savepoint_range: Optional[SavePointsRange] = None + self.id = -1 + + if bindings is None: + assert isinstance( + task, SavePoint + ), f"A pipeline node with no bindings can only be a Savepoint, got {task}" + self.bindings = [] + else: + assert isinstance(task, Pipe), f"Bindings are only allowed for Pipe tasks, got {task}" + if len(bindings) != len(task.arguments): + raise ValueError( + f"Expected {len(task.arguments)} bindings but got " + f"{len(bindings)} for task {task}. arguments: " + f"{task.arguments}, bindings: {bindings}" + ) + for bind, arg in zip(bindings, task.arguments): + if bind.container_type != arg.container_type: + raise TypeError( + f"Binding {bind} is not compatible with argument {arg} " f"for task {task}" + ) + # These maps the declarations of the pipeline to the task arguments declarations + self.bindings = list(bindings) + assert len(self.bindings) == len(task.arguments) + + @property + def arguments(self) -> list[ContainerDeclaration]: + if isinstance(self.task, SavePoint): + # SavePoints do not have bindings, so we return the task arguments directly + return [ + ContainerDeclaration( + name=x.name, + container_type=x.container_type, + ) + for x in self.task.arguments + ] + elif isinstance(self.task, Pipe): + # For Pipes, we return the bindings, which are the pipeline declarations + # that map to the task arguments + return self.bindings + else: + raise TypeError(f"Unsupported task type: {type(self.task)}") + + def add_successor(self, node: PipelineNode) -> PipelineNode: + node.predecessors.append(self) + self.successors.append(node) + return node + + def sorted_successors(self) -> list[PipelineNode]: + """ + Return the successors of this node sorted by their deterministic sort key. + This is needed to ensure that the savepoint ranges are assigned in a + deterministic order, otherwise the caches could cause problems. + """ + return sorted(self.successors, key=deterministic_sort_key) + + def prerequisites_for( + self, + model: ReadOnlyModel, + requests: Requests, + pipeline_configuration: PipelineConfiguration, + storage_provider: StorageProvider, + ) -> Requests: + # In order to execute the task we might need to remap the requests + # from the pipeline declarations to the task arguments declarations. + # So we apply a remap before and after + if isinstance(self.task, SavePoint): + # SavePoints have a different semantics for prerequisites, so we call the method + # from the SavePoint class + assert ( + self.savepoint_range is not None + ), "SavePoint range must be set before calling prerequisites_for on a SavePoint" + return self.task.prerequisites_for( + requests=requests, + configuration_id=self.configuration_id(pipeline_configuration), + storage_provider=storage_provider, + savepoint_range=self.savepoint_range, + ) + elif isinstance(self.task, Pipe): + # Use the bindings to map the requests to the pipe arguments + pipe_requests: list[ObjectSet] = [ + requests.get(decl, ObjectSet(decl.container_type.kind, set())) + for decl in self.bindings + ] + # Ask the pipe for its prerequisites + outgoing = self.task.prerequisites_for( + model=model, + requests=pipe_requests, + ) + # Map back the requests to the pipeline declarations + result = requests.clone() + for decl, request in zip(self.bindings, outgoing): + result[decl] = request + return result + else: + raise TypeError(f"Unsupported task type: {type(self.task)}") + + def configuration_id(self, configuration: PipelineConfiguration) -> ConfigurationId: + hasher = sha256() + if isinstance(self.task, Pipe) and self.task in configuration: + hasher.update(configuration[self.task].encode()) + elif isinstance(self.task, SavePoint): + for pipe in sorted(self.pipe_dependencies, key=lambda x: (x.name, x.signature())): + hasher.update(configuration.get(pipe, "").encode()) + return hasher.hexdigest() + + def run( + self, + model: ReadOnlyModel, + containers: ContainerSet, + incoming: Requests, + outgoing: Requests, + pipeline_configuration: PipelineConfiguration, + storage_provider: StorageProvider, + ) -> ObjectDependencies | None: + """Forward the run call to the task, but remap the requests and containers.""" + if isinstance(self.task, SavePoint): + assert ( + self.savepoint_range is not None + ), "SavePoint range must be set before calling run on a SavePoint" + self.task.run( + containers=containers, + incoming=incoming, + outgoing=outgoing, + configuration_id=self.configuration_id(pipeline_configuration), + storage_provider=storage_provider, + savepoint_range=self.savepoint_range, + ) + # SavePoints do not return any dependencies + return None + elif isinstance(self.task, Pipe): + pipe_containers = [containers[decl] for decl in self.bindings] + pipe_incoming = [ + incoming.get(decl, ObjectSet(decl.container_type.kind, set())) + for decl in self.bindings + ] + pipe_outgoing = [ + outgoing.get(decl, ObjectSet(decl.container_type.kind, set())) + for decl in self.bindings + ] + deps = self.task.run( + model=model, + containers=pipe_containers, + incoming=pipe_incoming, + outgoing=pipe_outgoing, + configuration=pipeline_configuration.get(self.task, ""), + ) + result: ObjectDependencies = [] + for index, index_deps in enumerate(deps): + container_type = self.task.signature()[index] + if container_type.access == TaskArgumentAccess.READ: + assert len(index_deps) == 0, ( + "An read only container cannot produce new objects so it can't add " + f"dependencies. For container {container_type.name} got dependencies " + f"{index_deps}" + ) + result.extend((self.bindings[index].name, obj, path) for obj, path in index_deps) + return result + else: + raise TypeError(f"Unsupported task type: {type(self.task)}") + + def __repr__(self): + return self.task.__repr__() diff --git a/python/revng/pypeline/pipeline_parser.py b/python/revng/pypeline/pipeline_parser.py new file mode 100644 index 000000000..d271f44ca --- /dev/null +++ b/python/revng/pypeline/pipeline_parser.py @@ -0,0 +1,378 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Optional + +import jsonschema +import yaml + +from .analysis import Analysis, AnalysisBinding +from .container import Container, ContainerDeclaration +from .pipeline import Artifact, Pipeline +from .pipeline_node import PipelineNode +from .task.pipe import Pipe +from .task.savepoint import SavePoint +from .utils.registry import get_registry + + +@dataclass(slots=True) +class Node: + """ + We use this to represent the graph of branches in the pipeline, + so we can compute the loading order of the branches. + """ + + content: Any + is_root: bool = True + successors: set[str] = field(default_factory=set) + + +def parse_pipe( + task: dict[str, Any], + container_decls: dict[str, ContainerDeclaration], +): + """Parse a pipe task from the JSON value.""" + pipes = get_registry(Pipe) # type: ignore[type-abstract] + pipe_name = task["pipe"] + if pipe_name not in pipes: + raise ValueError( + f"Pipe {pipe_name} is not registered, available pipes: " f"{sorted(pipes.keys())}" + ) + pipe_ty = pipes[pipe_name] + + pipe_args = task.get("arguments", []) + bindings = [] + for arg in pipe_args: + if arg not in container_decls: + raise ValueError( + f"While parsing {pipe_name}'s arguments found container `{arg}` " + "that is not declared in the pipeline" + ) + bindings.append(container_decls[arg]) + + name = task.get("name") + configuration = task.get("configuration", "") + + return PipelineNode( + task=pipe_ty( + name=name, + static_configuration=configuration, + ), + bindings=bindings, + ) + + +def parse_savepoint( + task: dict[str, Any], + container_decls: dict[str, ContainerDeclaration], +): + """Parse a savepoint task from the JSON value.""" + name = task["savepoint"]["name"] + containers = task["savepoint"]["containers"] + args = [] + for container_name in containers: + if container_name not in container_decls: + raise ValueError(f"Container {container_name} is not declared in the pipeline") + args.append(container_decls[container_name]) + return PipelineNode(SavePoint(name=name, to_save=args)) + + +def parse_artifacts( + node_artifacts: list[Any], + artifacts: set[Artifact], + target_node: PipelineNode, + container_decls: dict[str, ContainerDeclaration], +): + """ + Parse artifacts from the node artifacts list and populate the artifacts dictionary. + """ + for artifact in node_artifacts: + name = artifact["name"] + container = artifact["container"] + if container not in container_decls: + raise ValueError( + f"Artifact {name} references container {container} that is not " + "declared in the pipeline" + ) + artifacts.add( + Artifact( + name=name, + node=target_node, + container=container_decls[container], + ) + ) + + +def parse_analyses( + node_analyses: list[Any], + analyses: set[AnalysisBinding], + target_node: PipelineNode, + container_decls: dict[str, ContainerDeclaration], +): + """ + Parse analyses from the node analyses list and populate the analyses dictionary. + """ + for analysis in node_analyses: + analysis_name = analysis["analysis"] + if analysis_name in analyses: + raise ValueError(f"Analysis {analysis_name} is defined multiple times in the pipeline") + containers = analysis["containers"] + bindings = [] + for container in containers: + if container not in container_decls: + raise ValueError( + f"Analysis {analysis_name} references container {container} that is " + "not declared in the pipeline" + ) + bindings.append(container_decls[container]) + + analyses_registry = get_registry(Analysis) # type: ignore[type-abstract] + if analysis_name not in analyses_registry: + raise ValueError( + f"Analysis {analysis_name} is not registered, available analyses: " + f"{sorted(analyses_registry.keys())}" + ) + analysis_ty: type[Analysis] = analyses_registry[analysis_name] + + name = analysis.get("name", analysis_name) + + analyses.add( + AnalysisBinding( + analysis=analysis_ty(name), + bindings=tuple(bindings), + node=target_node, + ) + ) + + +def parse_task( + task: Any, + artifacts: set[Artifact], + analyses: set[AnalysisBinding], + container_decls: dict[str, ContainerDeclaration], + parent: Optional[PipelineNode] = None, +) -> PipelineNode: + """ + Parse a single task from the JSON value. + The JSON value should contain a dictionary with the branch content. + """ + + # Create the PipelineNode + res: PipelineNode + if "savepoint" in task: + res = parse_savepoint( + task=task, + container_decls=container_decls, + ) + else: + res = parse_pipe( + task=task, + container_decls=container_decls, + ) + + # Parse artifacts + node_artifacts = task.get("artifacts", []) + parse_artifacts( + node_artifacts=node_artifacts, + artifacts=artifacts, + target_node=res, + container_decls=container_decls, + ) + + # Parse analyses + node_analyses = task.get("analyses", []) + parse_analyses( + node_analyses=node_analyses, + analyses=analyses, + target_node=res, + container_decls=container_decls, + ) + + # Connect + if parent is not None: + parent.add_successor(res) + return res + + +def parse_branch( + node: Node, + graph: dict[str, PipelineNode], + artifacts: set[Artifact], + analyses: set[AnalysisBinding], + container_decls: dict[str, ContainerDeclaration], +) -> PipelineNode: + """ + Parse the branch and return the last node in the branch. + """ + parent: Optional[PipelineNode] = None + if "from" in node.content: + parent_name = node.content["from"] + if parent_name not in graph: + raise ValueError(f"Branch {parent_name} is not defined in the pipeline") + parent = graph[parent_name] + + # Parse the nodes + tasks = node.content.get("tasks", []) + for task in tasks: + if "pipe" not in task and "savepoint" not in task: + raise ValueError("Task must have either a pipe or a savepoint") + if "pipe" in task and "savepoint" in task: + raise ValueError("Task cannot have both a pipe and a savepoint") + + # Parse the node + res = parse_task( + task=task, + artifacts=artifacts, + analyses=analyses, + container_decls=container_decls, + parent=parent, + ) + parent = res + + assert parent is not None, "A branch cannot be empty and not have a parent" + return parent + + +def parse_branches( + branches: dict[str, Any], + artifacts: set[Artifact], + analyses: set[AnalysisBinding], + container_decls: dict[str, ContainerDeclaration], +) -> PipelineNode: + """ + Parse the branches from the JSON value. + The JSON value should contain a dictionary of branches with their names and tasks. + """ + + # Find the root branch + graph: dict[str, Node] = {} + for name, branch in branches.items(): + node = Node(content=branch) + graph[name] = node + if "from" in branch: + node.is_root = False + from_node = graph[branch["from"]] + from_node.successors.add(name) + + # Check that there is only one root branch + roots = sorted(name for name, node in graph.items() if node.is_root) + if len(roots) != 1: + raise ValueError(f"There should be exactly one root branch but found {roots}") + root = roots[0] + + # Parse the branches in the correct order (DFS) + result: dict[str, PipelineNode] = {} + queue: list[str] = [root] + root_node: PipelineNode | None = None + while queue: + name = queue.pop(0) + node = graph[name] + # Parse the branch + result[name] = parse_branch( + node=node, + graph=result, + container_decls=container_decls, + artifacts=artifacts, + analyses=analyses, + ) + if root_node is None: + root_node = result[name] + # Result contains the END of each branch, so we need to find the root node + # by following the predecessors + while root_node.predecessors: + root_node = root_node.predecessors[0] + # Enqueue the successors + for successor in node.successors: + assert successor not in result, ( + f"The pipeline has to be a tree, not a DAG. Branch {successor} " + "is defined multiple times" + ) + queue.append(successor) + # Return the root node + assert root_node is not None, "The root node should be defined" + return root_node + + +def parse_container_decls( + containers: list[Any], +) -> dict[str, ContainerDeclaration]: + """ + Parse the container declarations from the JSON value. + The JSON value should contain a list of containers with their names and types. + """ + containers_registry = get_registry(Container) # type: ignore[type-abstract] + container_decls: dict[str, ContainerDeclaration] = {} + for container in containers: + name = container["name"] + ty = container["type"] + if ty not in containers_registry: + raise ValueError( + f"Container type {ty} is not registered, the available types " + "are: {list(sorted(containers_registry.keys()))}" + ) + + container_decls[name] = ContainerDeclaration( + name=name, + container_type=containers_registry[ty], + ) + return container_decls + + +def schema() -> dict[str, Any]: + """ + Return the jsonschema for the pipeline. + """ + root = Path(__file__).resolve().parent + with open(root / "pipeline_schema.yml", "r", encoding="utf-8") as f: + return yaml.safe_load(f) + + +def load_pipeline(values: Any) -> Pipeline: + """ + Load a pipeline from parsed JSON / YAML / TOML. + """ + validator = jsonschema.Draft7Validator(schema()) + validator.validate(values) + # Yeah the validator already checks that everything is correct, but + # mypy doesn't + + # Parse create all the container declarations + container_decls: dict[str, ContainerDeclaration] = parse_container_decls(values["containers"]) + + # These will get filled while parsing branches + artifacts: set[Artifact] = set() + analyses: set[AnalysisBinding] = set() + + root: PipelineNode = parse_branches( + branches=values["branches"], + container_decls=container_decls, + artifacts=artifacts, + analyses=analyses, + ) + return Pipeline( + declarations=set(container_decls.values()), + root=root, + artifacts=artifacts, + analyses=analyses, + ) + + +def load_pipeline_yaml(yaml_data: str) -> Pipeline: + """ + Load a pipeline from a YAML string. + """ + values = yaml.safe_load(yaml_data) + return load_pipeline(values) + + +def load_pipeline_yaml_file(file: str) -> Pipeline: + """ + Load a pipeline from a YAML file.""" + with open(file, "r", encoding="utf-8") as f: + values = yaml.safe_load(f) + return load_pipeline(values) diff --git a/python/revng/pypeline/pipeline_schema.yml b/python/revng/pypeline/pipeline_schema.yml new file mode 100644 index 000000000..4685faf1c --- /dev/null +++ b/python/revng/pypeline/pipeline_schema.yml @@ -0,0 +1,150 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +$schema: http://json-schema.org/draft-07/schema# +title: Pipeline Schema +type: object +required: + - containers + - branches +additionalProperties: false +properties: + containers: + type: array + uniqueItems: true + items: + $ref: "#/$defs/container_decl" + branches: + type: object + required: + - root + patternProperties: + ^[a-zA-Z0-9_]+$: + type: object + required: + - tasks + additionalProperties: false + properties: + from: + type: string + description: The current branch will start from the end of the given branch. + example: root + tasks: + type: array + items: + anyOf: + - $ref: "#/$defs/pipe" + - type: object + required: + - savepoint + additionalProperties: false + properties: + savepoint: + $ref: "#/$defs/savepoint" + description: A list of tasks in the branch, each task is a pipe. + minItems: 1 +$defs: + container_decl: + type: object + required: + - name + - type + additionalProperties: false + properties: + name: + type: string + description: The name of the container. + example: my_container + type: + type: string + description: The type of the container. + example: llvm_module + args: + type: array + items: + type: string + description: A list of container names that cannot be empty. + example: + - container1 + - container2 + minItems: 1 + pipe: + type: object + additionalProperties: false + required: + - pipe + - arguments + properties: + name: + type: string + description: The name of the pipe task. + example: my_pipe_task + pipe: + type: string + description: The name of the pipe. + example: my_pipe + configuration: + type: string + description: The static configuration for the pipe, if any. + default: "" + example: intel_asm=1 + arguments: + $ref: "#/$defs/args" + description: The names of the containers to pass to the pipe. + artifacts: + type: array + items: + $ref: "#/$defs/artifact" + minItems: 1 + analyses: + type: array + items: + $ref: "#/$defs/analysis" + minItems: 1 + savepoint: + type: object + required: + - name + - containers + additionalProperties: false + properties: + name: + type: string + description: The name of the savepoint. + containers: + $ref: "#/$defs/args" + description: The names of the containers the savepoint will cache. + analysis: + type: object + required: + - analysis + - containers + additionalProperties: false + properties: + name: + type: string + description: The name that can be used from the CLI to refer to the analysis + in this point of the pipeline. If not passed it defaults to the name of + the analysis. + analysis: + type: string + description: The name of the analysis to run. + containers: + $ref: "#/$defs/args" + description: The names of the containers the analysis will need. + artifact: + type: object + required: + - name + - container + additionalProperties: false + properties: + name: + type: string + description: The name of the artifact. + example: my_artifact + container: + type: string + description: The container that will produce the artifact. + example: my_container diff --git a/python/revng/pypeline/pyproject.toml b/python/revng/pypeline/pyproject.toml new file mode 100644 index 000000000..74d9dfbb1 --- /dev/null +++ b/python/revng/pypeline/pyproject.toml @@ -0,0 +1,26 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +[project] +name = "pypeline" +dynamic = ["version", "readme"] +dependencies = [ + "click==8.2.1", + "jsonschema", + "pyyaml", +] + +[project.scripts] +pype = "revng.pypeline.main:run" + +[tool.setuptools.package-data] +pypeline = ["revng/pypeline/pipeline_schema.yml"] + +[build-system] +requires = ["setuptools", "wheel"] +build-backend = "setuptools.build_meta" + +[tool.setuptools.dynamic] +version = {attr = "revng.pypeline.__version__"} +readme = {file = ["README.md"], content-type = "text/markdown"} diff --git a/python/revng/pypeline/schedule/__init__.py b/python/revng/pypeline/schedule/__init__.py new file mode 100644 index 000000000..79d95cc20 --- /dev/null +++ b/python/revng/pypeline/schedule/__init__.py @@ -0,0 +1,3 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# diff --git a/python/revng/pypeline/schedule/schedule.py b/python/revng/pypeline/schedule/schedule.py new file mode 100644 index 000000000..a3614e8c8 --- /dev/null +++ b/python/revng/pypeline/schedule/schedule.py @@ -0,0 +1,239 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +import logging +from graphlib import TopologicalSorter +from typing import Any, Dict, List, Optional, Set + +import yaml + +from revng.pypeline.container import ConfigurationId, ContainerDeclaration, ContainerSet +from revng.pypeline.graph import Graph +from revng.pypeline.model import ReadOnlyModel +from revng.pypeline.pipeline_node import PipelineConfiguration +from revng.pypeline.storage.storage_provider import StorageProvider +from revng.pypeline.task.pipe import Pipe +from revng.pypeline.task.requests import Requests +from revng.pypeline.task.savepoint import SavePoint +from revng.pypeline.task.task import ObjectDependencies, TaskArgumentAccess + +from .scheduled_task import ScheduledTask + +logger = logging.getLogger(__name__) + + +class Schedule: + """ + A pipeline, given a set of requests, a model and the configuration of pipes, can produce a + Schedule, i.e., a list (actually, a DAG) of tasks that need to be run with a certain set of + requests in order to fulfill the requests. + """ + + def __init__( + self, + declarations: Set[ContainerDeclaration], + root: ScheduledTask, + pipeline_configuration: PipelineConfiguration, + ): + self.declarations = set(declarations) + self.root = root + self.tasks: Set[ScheduledTask] = set(root.all_dependencies()) + self.pipeline_configuration: PipelineConfiguration = pipeline_configuration + + def graph(self) -> Graph: + """ + Produce a graph for debugging purposes. + """ + + graph = Graph() + + nodes_map: Dict[ScheduledTask, Graph.Node] = {} + + def label(requests: Requests) -> str: + result = "" + for container, objects in requests.items(): + result += f"{container.name}:\n " + "\n ".join(str(x) for x in objects) + "\n" + return result + + def get_node(node: ScheduledTask) -> Graph.Node: + if node not in nodes_map: + new_node = Graph.Node(node.node.task.name) + new_node.completed = node.completed + for argument in node.node.arguments: + new_node.entries.append(argument.name) + + nodes_map[node] = new_node + graph.nodes.add(new_node) + + return nodes_map[node] + + to_visit: List[ScheduledTask] = [self.root] + visited: Set[ScheduledTask] = set() + while to_visit: + node = to_visit.pop() + graph_node = get_node(node) + node_inputs: List[ContainerDeclaration] = list(node.node.arguments) + + for predecessor in node.depends_on: + for source_index, argument in enumerate(predecessor.node.task.arguments): + if argument.access == TaskArgumentAccess.READ or argument not in node_inputs: + continue + + destination_index = node_inputs.index(argument) + + source_node = get_node(predecessor) + new_edge = Graph.Edge( + source_node, + graph_node, + source_port=source_index, + destination_port=destination_index, + head_label=label(predecessor.outgoing), + tail_label=label(node.incoming), + ) + graph.edges.add(new_edge) + + if predecessor not in visited: + to_visit.append(predecessor) + visited.add(predecessor) + + return graph + + def run( + self, + model: ReadOnlyModel, + storage_provider: StorageProvider, + ) -> ContainerSet: + # Produce a set of working containers + working_containers: ContainerSet = { + declaration: declaration.instance() for declaration in self.declarations + } + + ready: ScheduledTask | None = self._pick_task() + + while ready: + logger.info("Running %s", ready.node.task.name) + + configuration: ConfigurationId = ready.node.configuration_id( + self.pipeline_configuration + ) + + task_dependencies: ObjectDependencies | None = ready.run( + model=model, + containers=working_containers, + pipeline_configuration=self.pipeline_configuration, + storage_provider=storage_provider, + ) + + for declaration, container in sorted( + working_containers.items(), key=lambda item: item[0].name + ): + logger.info(" %s: %s", declaration.name, str(container.objects())) + + if isinstance(ready.node.task, Pipe): + assert task_dependencies is not None + assert ( + ready.node.savepoint_range is not None + ), "Savepoint range should be set for all Pipes" + storage_provider.add_dependencies( + ready.node.savepoint_range, configuration, task_dependencies + ) + else: + assert task_dependencies is None + + ready = self._pick_task() + + return working_containers + + def _pick_task(self) -> Optional[ScheduledTask]: + # TODO: use a graph + for task in self.tasks: + if task.completed: + continue + + ready = True + for dependency in task.depends_on: + if not dependency.completed: + ready = False + break + + if ready: + return task + + return None + + def serialize(self) -> str: + """Serialize the Schedule to a YAML string""" + + containers = [] + for container in self.declarations: + containers.append({"name": container.name, "type": container.container_type.__name__}) + + toposorter: TopologicalSorter[ScheduledTask] = TopologicalSorter() + for task in self.tasks: + toposorter.add(task, *task.depends_on) + + tasks: list[Any] = [] + visited_tasks: list[ScheduledTask] = [] + for task in toposorter.static_order(): + if isinstance(task.node.task, Pipe): + pipe: Pipe = task.node.task + + args = [] + for declaration in task.node.bindings: + incoming = [x.serialize() for x in task.incoming.get(declaration, [])] + outgoing = [x.serialize() for x in task.outgoing.get(declaration, [])] + args.append( + {"name": declaration.name, "incoming": incoming, "outgoing": outgoing} + ) + + tasks.append( + { + "type": "Pipe", + "node_id": task.node.id, + "name": pipe.name, + "depends_on": [visited_tasks.index(t) for t in task.depends_on], + "static_config": pipe.static_configuration, + "dynamic_config": self.pipeline_configuration.get(pipe, ""), + "args": args, + } + ) + elif isinstance(task.node.task, SavePoint): + savepoint = task.node.task + + sp_containers = [] + for declaration in self.declarations: + incoming = [x.serialize() for x in task.incoming.get(declaration, [])] + outgoing = [x.serialize() for x in task.outgoing.get(declaration, [])] + if len(incoming) == 0 and len(outgoing) == 0: + continue + + sp_containers.append( + { + "name": declaration.name, + "configuration_hash": task.node.configuration_id( + self.pipeline_configuration + ), + "incoming": incoming, + "outgoing": outgoing, + } + ) + + assert task.node.savepoint_range is not None + tasks.append( + { + "type": "SavePoint", + "node_id": task.node.id, + "depends_on": [visited_tasks.index(t) for t in task.depends_on], + "name": savepoint.name, + "id": task.node.savepoint_range.start, + "containers": sp_containers, + } + ) + else: + raise ValueError(f"Unknown task: {type(task.node.task).__name__}") + visited_tasks.append(task) + + return yaml.safe_dump({"containers": containers, "tasks": tasks}) diff --git a/python/revng/pypeline/schedule/scheduled_task.py b/python/revng/pypeline/schedule/scheduled_task.py new file mode 100644 index 000000000..cf0b5f5dd --- /dev/null +++ b/python/revng/pypeline/schedule/scheduled_task.py @@ -0,0 +1,109 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Generator + +from revng.pypeline.container import ContainerSet +from revng.pypeline.model import ReadOnlyModel +from revng.pypeline.pipeline_node import PipelineConfiguration, PipelineNode +from revng.pypeline.storage.storage_provider import StorageProvider +from revng.pypeline.task.requests import Requests +from revng.pypeline.task.task import ObjectDependencies + + +@dataclass(slots=True) +class ScheduledTask: + """ + The scheduling works by figuring out which tasks to run, and which dependencies + each task can fulfill. + + The schedule is stored as a Directed Acyclic Graph (DAG) of ScheduledTask objects. + + While the schedule is originally a path of the pipeline tree, we can apply + transformations to it that can make it a DAG. An example of this is kind of + transformations is the parallelization of tasks, where a ScheduleTask is + split into multiple ScheduledTask objects that run in parallel the same task, + with partitioned inputs and outputs requests. + Moreover, using a DAG opens up the possibility of merging multiple schedules + into a single one to improve performance of batch jobs. + + As Tasks should be side-effect free, multiple ScheduledTasks can share + the same PipelineNode instance. + """ + + node: PipelineNode + """The node of the pipeline we scheduled.""" + + completed: bool = False + """This is used for debug, and asserting that a schedule is not run twice.""" + + outgoing: Requests = field(default_factory=Requests) + """These are the objects this task is supposed to compute and put in each container.""" + + incoming: Requests = field(default_factory=Requests) + """ + These are the dependencies of the task, i.e. the objects that this task + needs to run. This is computed during the scheduling phase, and is only used + to check that the task can run. + """ + + depends_on: list[ScheduledTask] = field(default_factory=list) + """ + These are the scheduled tasks that this task depends on, after the scheduling + phase this list will contain only 0 or 1 elements as it's a path, but + the schedule can become a DAG. + """ + + def request(self, incoming: Requests, outgoing: Requests) -> None: + """ + Add requests to the incoming and outgoing requests of this task. + """ + self.incoming.merge(incoming) + self.outgoing.merge(outgoing) + + def run( + self, + model: ReadOnlyModel, + containers: ContainerSet, + pipeline_configuration: PipelineConfiguration, + storage_provider: StorageProvider, + ) -> ObjectDependencies | None: + """ + Run the task with the requests computed during the scheduling phase. + """ + if self.completed: + raise RuntimeError(f"ScheduledTask {self.node} has already been run.") + self.incoming.check(containers) + result = self.node.run( + model=model, + containers=containers, + incoming=self.incoming, + outgoing=self.outgoing, + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + self.outgoing.check(containers) + self.completed = True + return result + + def all_dependencies(self) -> Generator[ScheduledTask, None, None]: + """ + From this task, yield all the transitive dependencies of this task. + This is used to get all the tasks involved in the execution of this task, + including the task itself. + """ + yield self + for dependency in self.depends_on: + yield from dependency.all_dependencies() + + def __hash__(self): + """ + The intended use of this hash is to be used to do visits on the schedule graph. + So this hash is just the id of the object, which is unique for each instance. + Two instances of the same task will have different ids, and thus different hashes. + """ + return hash(id(self)) diff --git a/python/revng/pypeline/storage/__init__.py b/python/revng/pypeline/storage/__init__.py new file mode 100644 index 000000000..c60869fd4 --- /dev/null +++ b/python/revng/pypeline/storage/__init__.py @@ -0,0 +1,48 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from pathlib import Path +from typing import Optional + +from revng.pypeline.model import Model +from revng.pypeline.utils.registry import get_singleton + +from .memory import InMemoryStorageProvider +from .null import NullStorageProvider +from .sqlite3 import SQlite3StorageProvider +from .storage_provider import StorageProvider + + +def storage_provider_factory( + model_path: str, + storage_url: Optional[str] = None, +) -> StorageProvider: + """ + Given a storage_url, return the appropriately instantiated StorageProvider. + If storage_url is None, it defaults to "sqlite://pypeline.db". + """ + storage_url = storage_url or "sqlite://pypeline.db" + resolved_model_path = Path(model_path).resolve() + + if storage_url.startswith("sqlite://"): + # Create an empty file if it does not exist + if not resolved_model_path.exists(): + model_ty: type[Model] = get_singleton(Model) # type: ignore[type-abstract] + empty_model = model_ty() + with resolved_model_path.open("wb") as f: + f.write(empty_model.serialize()) + + return SQlite3StorageProvider( + db_path=storage_url[len("sqlite://") :], + model_path=resolved_model_path, + ) + elif storage_url.startswith("memory"): + return InMemoryStorageProvider() + elif storage_url.startswith("null"): + return NullStorageProvider() + else: + raise ValueError( + "Unknown storage provider `%s`. " + "Please set the PYPELINE_STORAGE environment variable." % storage_url, + ) diff --git a/python/revng/pypeline/storage/memory.py b/python/revng/pypeline/storage/memory.py new file mode 100644 index 000000000..fef2063df --- /dev/null +++ b/python/revng/pypeline/storage/memory.py @@ -0,0 +1,134 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +from collections import defaultdict +from dataclasses import dataclass +from datetime import datetime +from typing import Iterable, Mapping + +from revng.pypeline.container import ContainerID +from revng.pypeline.model import ModelPath, ModelPathSet +from revng.pypeline.object import ObjectID +from revng.pypeline.task.task import ObjectDependencies + +from .storage_provider import ConfigurationId, ContainerLocation, ProjectMetadata, SavepointID +from .storage_provider import SavePointsRange, StorageProvider +from .util import _REVNG_VERSION_PLACEHOLDER + + +@dataclass +class DependencyEntry: + savepoint_start: SavepointID + savepoint_end: SavepointID + container_id: ContainerID + configuration_id: ConfigurationId + object_id: ObjectID + + +class InMemoryStorageProvider(StorageProvider): + """A simple in-memory storage provider for testing purposes. + This is not thread-safe and should not be used in production. + """ + + def __init__(self): + self.model = b"" + self.storage: dict[ContainerLocation, dict[ObjectID, bytes]] = {} + self.dependencies: dict[str, list[DependencyEntry]] = defaultdict(list) + self.last_change = datetime.now() + + def has( + self, + location: ContainerLocation, + keys: Iterable[ObjectID], + ) -> Iterable[ObjectID]: + if location not in self.storage: + return [] + storage = self.storage[location] + return [key for key in keys if key in storage] + + def get( + self, + location: ContainerLocation, + keys: Iterable[ObjectID], + ) -> Mapping[ObjectID, bytes]: + if location not in self.storage: + raise KeyError(f"Savepoint {location} not found in storage.") + storage = self.storage[location] + return {k: storage[k] for k in keys if k in storage} + + def add_dependencies( + self, + savepoint_range: SavePointsRange, + configuration_id: ConfigurationId, + deps: ObjectDependencies, + ) -> None: + for container_name, obj, path in deps: + self.dependencies[path].append( + DependencyEntry( + savepoint_range.start, + savepoint_range.end, + container_name, + configuration_id, + obj, + ) + ) + self.last_change = datetime.now() + + def put( + self, + location: ContainerLocation, + values: Mapping[ObjectID, bytes], + ) -> None: + self.storage.setdefault(location, {}) + for key, value in values.items(): + self.storage[location][key] = value + self.last_change = datetime.now() + + def _invalidate(self, path: ModelPath) -> None: + for key, entries in self.dependencies.items(): + if key != path: + continue + + # TODO: this double loop is very inefficient + for entry in entries: + for container_loc, objects in self.storage.items(): + if ( + container_loc.savepoint_id < entry.savepoint_start + or container_loc.savepoint_id > entry.savepoint_end + or container_loc.container_id != entry.container_id + or container_loc.configuration_id != entry.configuration_id + ): + continue + if entry.object_id in objects: + del objects[entry.object_id] + + def invalidate(self, invalidation_list: ModelPathSet) -> None: + for path in invalidation_list: + self._invalidate(path) + self.last_change = datetime.now() + + def get_model(self) -> bytes: + return self.model + + def set_model(self, new_model: bytes): + self.model = new_model + self.last_change = datetime.now() + + def metadata(self) -> ProjectMetadata: + """ + Fetch metadata about the current project + """ + return ProjectMetadata( + last_change=self.last_change, + revng_version=_REVNG_VERSION_PLACEHOLDER, + ) + + def prune_objects(self): + """ + Prunes all the objects (except metadata) from storage + """ + self.storage.clear() + self.dependencies.clear() diff --git a/python/revng/pypeline/storage/null.py b/python/revng/pypeline/storage/null.py new file mode 100644 index 000000000..239ef2f3e --- /dev/null +++ b/python/revng/pypeline/storage/null.py @@ -0,0 +1,81 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +from datetime import datetime +from typing import Iterable, Mapping + +from revng.pypeline.container import ConfigurationId +from revng.pypeline.model import ModelPathSet +from revng.pypeline.object import ObjectID +from revng.pypeline.task.task import ObjectDependencies + +from .storage_provider import ContainerLocation, ProjectMetadata, SavePointsRange, StorageProvider +from .util import _REVNG_VERSION_PLACEHOLDER + + +class NullStorageProvider(StorageProvider): + """The /dev/null of storage providers. It stores nothing and caches nothing. + It just keeps track of the model, project ID, and last change time as those + are required by the interface. + """ + + def __init__(self): + self.model = b"" + self.last_change = datetime.now() + + def has( + self, + location: ContainerLocation, + keys: Iterable[ObjectID], + ) -> Iterable[ObjectID]: + return [] + + def get( + self, + location: ContainerLocation, + keys: Iterable[ObjectID], + ) -> Mapping[ObjectID, bytes]: + assert not keys, "NullStorageProvider does not support get operation." + return {} + + def add_dependencies( + self, + savepoint_range: SavePointsRange, + configuration_id: ConfigurationId, + deps: ObjectDependencies, + ) -> None: + self.last_change = datetime.now() + + def put( + self, + location: ContainerLocation, + values: Mapping[ObjectID, bytes], + ) -> None: + self.last_change = datetime.now() + + def invalidate(self, invalidation_list: ModelPathSet) -> None: + self.last_change = datetime.now() + + def get_model(self) -> bytes: + return self.model + + def set_model(self, new_model: bytes): + self.model = new_model + self.last_change = datetime.now() + + def metadata(self) -> ProjectMetadata: + """ + Fetch metadata about the current project + """ + return ProjectMetadata( + last_change=self.last_change, + revng_version=_REVNG_VERSION_PLACEHOLDER, + ) + + def prune_objects(self): + """ + Prunes all the objects (except metadata) from storage + """ diff --git a/python/revng/pypeline/storage/sqlite3.py b/python/revng/pypeline/storage/sqlite3.py new file mode 100644 index 000000000..fcca85205 --- /dev/null +++ b/python/revng/pypeline/storage/sqlite3.py @@ -0,0 +1,242 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +import sqlite3 +from datetime import datetime, timezone +from pathlib import Path +from typing import Iterable, Mapping + +from revng.pypeline.container import ConfigurationId +from revng.pypeline.model import ModelPathSet +from revng.pypeline.object import ObjectID +from revng.pypeline.task.task import ObjectDependencies +from revng.pypeline.utils.registry import get_singleton + +from .storage_provider import ContainerLocation, ProjectMetadata, SavePointsRange, StorageProvider +from .util import _REVNG_VERSION_PLACEHOLDER + +CREATE_TABLES = """ +CREATE TABLE IF NOT EXISTS project( + id TEXT PRIMARY KEY CHECK (id = 0), + last_change REAL, + revng_version TEXT +) STRICT; + +CREATE TABLE IF NOT EXISTS objects( + savepoint_id INT NOT NULL, + container_id TEXT NOT NULL, + configuration_hash TEXT NOT NULL, + object_id TEXT NOT NULL, + content BLOB NOT NULL, + PRIMARY KEY (savepoint_id, container_id, configuration_hash, object_id) +) STRICT; + +CREATE INDEX IF NOT EXISTS savepoint_id_on_object ON objects(savepoint_id); +CREATE INDEX IF NOT EXISTS container_id_on_object ON objects(container_id); +CREATE INDEX IF NOT EXISTS configuration_hash_on_object ON objects(configuration_hash); +CREATE INDEX IF NOT EXISTS object_id_on_object ON objects(object_id); + +CREATE TABLE IF NOT EXISTS dependencies( + savepoint_id_start INT NOT NULL, + savepoint_id_end INT NOT NULL, + container_id TEXT NOT NULL, + configuration_hash TEXT NOT NULL, + object_id TEXT NOT NULL, + model_path_hash TEXT NOT NULL, + PRIMARY KEY (savepoint_id_start, savepoint_id_end, container_id, + configuration_hash, object_id, model_path_hash) +) STRICT; + +CREATE INDEX IF NOT EXISTS model_path_hash_on_dependencies ON dependencies(model_path_hash); +""" + +HAS_QUERY = """ +SELECT object_id FROM objects +WHERE + savepoint_id = ? + AND container_id = ? + AND configuration_hash = ? + AND object_id IN ({id_list})""" + +GET_QUERY = """ +SELECT object_id, content FROM objects +WHERE + savepoint_id = ? + AND container_id = ? + AND configuration_hash = ? + AND object_id IN ({id_list})""" + +PUT_QUERY = """ +REPLACE INTO objects(savepoint_id, container_id, configuration_hash, object_id, content) +VALUES (?, ?, ?, ?, ?) +""" + +INVALIDATE_QUERY = """ +DELETE FROM objects +WHERE rowid IN ( + SELECT objects.rowid + FROM objects + JOIN dependencies + WHERE dependencies.model_path_hash IN ({model_paths}) + AND objects.container_id = dependencies.container_id + AND objects.configuration_hash = dependencies.configuration_hash + AND objects.savepoint_id >= dependencies.savepoint_id_start + AND objects.savepoint_id <= dependencies.savepoint_id_end +); + +DELETE FROM dependencies +WHERE dependencies.model_path_hash IN ({model_paths}); +""" + + +class CursorWrapper: + def __init__(self, connection: sqlite3.Connection): + self.connection = connection + self.cursor: sqlite3.Cursor | None = None + + def __enter__(self) -> sqlite3.Cursor: + assert self.cursor is None + self.cursor = self.connection.cursor() + return self.cursor + + def __exit__(self, exc_type, exc_val, exc_tb): + assert self.cursor is not None + if exc_type is not None: + self.connection.rollback() + else: + self.connection.commit() + self.cursor.close() + return False + + +class SQlite3StorageProvider(StorageProvider): + """StorageProvider implementation with backing sqlite3 db""" + + def __init__(self, db_path: str, model_path: str | Path): + self._model_path = Path(model_path) + self._connection = sqlite3.connect(db_path, autocommit=False) + self._connection.commit() + self._init_tables() + + def _cursor(self) -> CursorWrapper: + return CursorWrapper(self._connection) + + def _init_tables(self): + with self._cursor() as cursor: + cursor.executescript(CREATE_TABLES) + + def _write_metadata(self, cursor: sqlite3.Cursor): + cursor.execute( + "REPLACE INTO project VALUES (0, ?, ?)", + (datetime.now().timestamp(), _REVNG_VERSION_PLACEHOLDER), + ) + + def has( + self, + location: ContainerLocation, + keys: Iterable[ObjectID], + ) -> Iterable[ObjectID]: + # NOTE: possible SQL injection, sqlite has a limit on parameters. If it + # can't be avoided chunk selects by 999 values + id_list = ",".join([f"'{key!s}'" for key in keys]) + with self._cursor() as cursor: + cursor.execute( + HAS_QUERY.format(id_list=id_list), + (location.savepoint_id, location.container_id, location.configuration_id), + ) + result = cursor.fetchall() + obj_id_ty: type[ObjectID] = get_singleton(ObjectID) # type: ignore[type-abstract] + return [obj_id_ty.deserialize(x[0]) for x in result] + + def get( + self, + location: ContainerLocation, + keys: Iterable[ObjectID], + ) -> Mapping[ObjectID, bytes]: + # NOTE: possible SQL injection, sqlite has a limit on parameters. If it + # can't be avoided chunk selects by 999 values + id_list = ",".join([f"'{key!s}'" for key in keys]) + with self._cursor() as cursor: + cursor.execute( + GET_QUERY.format(id_list=id_list), + (location.savepoint_id, location.container_id, location.configuration_id), + ) + result = cursor.fetchall() + obj_id_ty: type[ObjectID] = get_singleton(ObjectID) # type: ignore[type-abstract] + return {obj_id_ty.deserialize(x[0]): x[1] for x in result} + + def add_dependencies( + self, + savepoint_range: SavePointsRange, + configuration_id: ConfigurationId, + deps: ObjectDependencies, + ) -> None: + with self._cursor() as cursor: + for container_id, object_id, model_path in deps: + cursor.execute( + "REPLACE INTO dependencies VALUES (?, ?, ?, ?, ?, ?)", + ( + savepoint_range.start, + savepoint_range.end, + container_id, + configuration_id, + object_id.serialize(), + model_path, + ), + ) + self._write_metadata(cursor) + + def put( + self, + location: ContainerLocation, + values: Mapping[ObjectID, bytes], + ) -> None: + with self._cursor() as cursor: + for object_id, content in values.items(): + cursor.execute( + PUT_QUERY, + ( + location.savepoint_id, + location.container_id, + location.configuration_id, + object_id.serialize(), + content, + ), + ) + self._write_metadata(cursor) + + def invalidate(self, invalidation_list: ModelPathSet) -> None: + with self._cursor() as cursor: + cursor.executescript( + INVALIDATE_QUERY.format( + model_paths=",".join(f"'{path}'" for path in invalidation_list) + ) + ) + self._write_metadata(cursor) + + def prune_objects(self): + with self._cursor() as cursor: + cursor.execute("DELETE FROM objects") + cursor.execute("DELETE FROM dependencies") + self._write_metadata(cursor) + + def get_model(self) -> bytes: + return self._model_path.read_bytes() + + def set_model(self, new_model: bytes): + self._model_path.write_bytes(new_model) + with self._cursor() as cursor: + self._write_metadata(cursor) + + def metadata(self) -> ProjectMetadata: + with self._cursor() as cursor: + cursor.execute("SELECT last_change, revng_version FROM project WHERE id is NULL") + result = cursor.fetchone() + + return ProjectMetadata( + last_change=datetime.fromtimestamp(result[0], timezone.utc), + revng_version=result[1], + ) diff --git a/python/revng/pypeline/storage/storage_provider.py b/python/revng/pypeline/storage/storage_provider.py new file mode 100644 index 000000000..d02a4e069 --- /dev/null +++ b/python/revng/pypeline/storage/storage_provider.py @@ -0,0 +1,158 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +from abc import ABC, abstractmethod +from dataclasses import dataclass +from datetime import datetime +from typing import Annotated, Iterable, Mapping + +from revng.pypeline.container import ConfigurationId, ContainerID +from revng.pypeline.model import ModelPathSet +from revng.pypeline.object import ObjectID +from revng.pypeline.task.task import ObjectDependencies + +SavepointID = Annotated[ + int, + """ + An integer that represents a savepoint, we want this to be an integer so we + can assign them doing a DFS traversal of the pipeline, which in turns allows + to efficiently represent a subtree of savepoints as a continuous range of + integers to avoid storing the dependencies multiple times. + """, +] + + +@dataclass(frozen=True, slots=True) +class ContainerLocation: + savepoint_id: SavepointID + container_id: ContainerID + configuration_id: ConfigurationId + + +@dataclass(frozen=True, slots=True) +class ProjectMetadata: + last_change: datetime + revng_version: str + + +@dataclass(frozen=True, slots=True) +class SavePointsRange: + """A range of savepoints, with inclusive extremes. + When the savepoints IDs are assigned in a DFS traversal of the pipeline, + any subtree of savepoints can be represented as a continuous range of integers, + avoiding the need to store the dependencies multiple times. + The starts are assigned with a pre-order traversal, and the ends are assigned + with a post-order traversal. The order is increase only when a savepoint is + visited for the first time. + """ + + start: SavepointID + """ + This is the smallest savepoint ID of the subtree of a node (root included). + Thus, if the node is a SavePoint, this is the ID of the SavePoint itself. + """ + + end: SavepointID + """This is the largest savepoint ID of the subtree of a node (root included). + Thus, if the node is a SavePoint, this is start + the number of savepoints + present in the subtree - 1 (minus the root itself).""" + + def __contains__(self, item: object) -> bool: + if not isinstance(item, int): + return False + return self.start <= item <= self.end + + def __len__(self) -> int: + """The number of savepoints in the range.""" + return self.end - self.start + 1 + + +class StorageProvider(ABC): + """This is the general interface for something that caches containers. + This can be in memory, on disk, in a database, etc. + This is a singleton and there should never be more than one instance of it. + """ + + @abstractmethod + def has( + self, + location: ContainerLocation, + keys: Iterable[ObjectID], + ) -> Iterable[ObjectID]: + """ + Get the available objects from the storage. + If the object is not found, it will not be included in the result. + """ + + @abstractmethod + def get( + self, + location: ContainerLocation, + keys: Iterable[ObjectID], + ) -> Mapping[ObjectID, bytes]: + """ + For each objects, return bytes that the container can ingest to + deserialize the object. + All the object **have to** be present, as one should call `get_available` + first to check it. + """ + + @abstractmethod + def add_dependencies( + self, + savepoint_range: SavePointsRange, + configuration_id: ConfigurationId, + deps: ObjectDependencies, + ) -> None: + """ + Store the dependencies between the objects and the model paths. + This is has to be called **BEFORE** `put`. + It can, and probably will, contain duplicated dependencies from previous + calls, but the storage should handle this gracefully. + """ + + @abstractmethod + def put( + self, + location: ContainerLocation, + values: Mapping[ObjectID, bytes], + ) -> None: + """ + Put a set of serialized objects into the storage. + This has always to be called **AFTER** `add_dependencies` + """ + + @abstractmethod + def invalidate(self, invalidation_list: ModelPathSet) -> None: + """ + Inform the storage that certain model paths are no longer valid. + The storage should use the stored dependencies to determine which objects + need to be invalidated. + """ + + @abstractmethod + def get_model(self) -> bytes: + """ + Get the model + """ + + @abstractmethod + def set_model(self, new_model: bytes): + """ + Set the model + """ + + @abstractmethod + def metadata(self) -> ProjectMetadata: + """ + Fetch metadata about the current project + """ + + @abstractmethod + def prune_objects(self): + """ + Prunes all the objects (except metadata) from storage + """ diff --git a/python/revng/pypeline/storage/util.py b/python/revng/pypeline/storage/util.py new file mode 100644 index 000000000..9eab73811 --- /dev/null +++ b/python/revng/pypeline/storage/util.py @@ -0,0 +1,5 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +_REVNG_VERSION_PLACEHOLDER = "1.0.0" diff --git a/python/revng/pypeline/task/__init__.py b/python/revng/pypeline/task/__init__.py new file mode 100644 index 000000000..79d95cc20 --- /dev/null +++ b/python/revng/pypeline/task/__init__.py @@ -0,0 +1,3 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# diff --git a/python/revng/pypeline/task/pipe.py b/python/revng/pypeline/task/pipe.py new file mode 100644 index 000000000..1f9430239 --- /dev/null +++ b/python/revng/pypeline/task/pipe.py @@ -0,0 +1,128 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from typing import Optional, final + +from revng.pypeline.container import Configuration, Container +from revng.pypeline.model import ReadOnlyModel +from revng.pypeline.object import ObjectSet +from revng.pypeline.utils.cabc import ABC, abstractmethod + +from .task import PipeObjectDependencies, TaskArgument, TaskArgumentAccess + + +class Pipe(ABC): + """ + A Pipe is a task that, given some input objects, a configuration string and the model, produces + some new objects. + """ + + @classmethod + @abstractmethod + def signature(cls) -> tuple[TaskArgument, ...]: + """ + While tasks can have a dynamic arguments, like a savepoint + can save different type of containers, a pipe has to have a static ones + that do not depend on the instance. This could be a class attribute, + but `@abstractclassattributes` does not exist in python. + """ + + @classmethod + def static_configuration_help(cls) -> Optional[str]: + """ + The text to display when the user asks for help on the + static configuration of this task. + Do not implement it, or return None, if the static configuration + argument should not be added. + """ + + __slots__: tuple = ("name", "static_configuration") + + def __init__( + self, + static_configuration: str = "", + name: str | None = None, + ): + self.name: str = name or self.__class__.__name__ + self.static_configuration: str = static_configuration + + @property + def arguments(self) -> list[TaskArgument]: + """ + Return the arguments of this pipe, which are the static configuration and the inputs and + outputs. + """ + return list(self.signature()) + + @property + def inputs(self) -> list[TaskArgument]: + """ + Return the inputs of this pipe, which are the arguments that are not outputs. + """ + return [arg for arg in self.signature() if arg.access != TaskArgumentAccess.WRITE] + + @property + def outputs(self) -> list[TaskArgument]: + """ + Return the outputs of this pipe, which are the arguments that are not inputs. + """ + return [arg for arg in self.signature() if arg.access != TaskArgumentAccess.READ] + + @final + def prerequisites_for( + self, + model: ReadOnlyModel, + requests: list[ObjectSet], + ) -> list[ObjectSet]: + """ + Given a set of requests, a configuration and a model, produce a new set + of requests that are required in order to run this pipeline successfully. + """ + # List of empty requests, one per argument + result = [ObjectSet(decl.container_type.kind, set()) for decl in self.arguments] + + # Cross-contaminate inputs and outputs + for idx, decl in enumerate(self.arguments): + # We must fill the readable containers + if decl.access == TaskArgumentAccess.WRITE: + continue + + for ridx, object_list in enumerate(requests): + # The requests must be for the writeable containers, + # otherwise it's not possible to satisfy them + if self.arguments[ridx].access == TaskArgumentAccess.READ: + assert len(object_list) == 0, ( + f"Expected an empty request for {self.arguments[ridx].name}, " + f"but got {object_list}." + ) + + result[idx].update( + model.move_to_kind( + object_list, + decl.container_type.kind, + ) + ) + + return result + + @abstractmethod + def run( + self, + model: ReadOnlyModel, + containers: list[Container], + incoming: list[ObjectSet], + outgoing: list[ObjectSet], + configuration: Configuration, + ) -> PipeObjectDependencies: + """ + Run the pipe with the given model. + The containers set is the set of ephemeral containers used for this run, + and they contains both the inputs and outputs of the pipe. + The incoming requests are the requests that were made to the pipe before + running it, they are mostly for validation purposes. + The outgoing requests are the objects that the pipe has to produce in + the requested containers as a result of running. + `containers`, `incoming`, and `outgoing` are all lists with the same + length as `SIGNATURE`. + """ diff --git a/python/revng/pypeline/task/requests.py b/python/revng/pypeline/task/requests.py new file mode 100644 index 000000000..7559d6bdc --- /dev/null +++ b/python/revng/pypeline/task/requests.py @@ -0,0 +1,106 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from __future__ import annotations + +from typing import Dict, MutableMapping, Optional, Set + +from revng.pypeline.container import ContainerDeclaration, ContainerSet +from revng.pypeline.object import ObjectSet + + +class Requests(MutableMapping[ContainerDeclaration, ObjectSet]): + """ + A map from a container declaration to the objects that we'd like to see there (an ObjectList). + """ + + def __init__(self, requests: Optional[Dict[ContainerDeclaration, ObjectSet]] = None): + self.requests: Dict[ContainerDeclaration, ObjectSet] = {} + if requests: + self.requests.update(requests) + + def clone(self) -> Requests: + result = Requests() + result.requests = {k: v.clone() for k, v in self.requests.items()} + return result + + def _requests_for(self, declaration: ContainerDeclaration) -> ObjectSet: + """ + If the key is not present, we create a new ObjectSet for it. + """ + kind = declaration.container_type.kind + if declaration in self.requests: + assert self.requests[declaration].kind == kind + else: + self.requests[declaration] = ObjectSet(kind=kind) + + return self.requests[declaration] + + def merge(self, other: Requests) -> None: + for container, objects in other.items(): + self._requests_for(container).update(objects) + + def extract(self, requested_containers: Set[ContainerDeclaration]) -> Requests: + results = Requests() + + for container, objects in list(self.requests.items()): + if container in requested_containers: + assert container in self.requests + results.requests[container] = objects + del self.requests[container] + + return results + + def check(self, containers: ContainerSet): + """ + Check if the given containers set satisfies these requests. + """ + # NOTE: this can be done either here or in the ContainerSet class, but + # currently it's a dict, so it's easier to do it here + for decl, objects in self.requests.items(): + if decl not in containers: + raise ValueError(f"Container {decl} is not present in the given ContainerSet.") + if not containers[decl].contains_all(objects): + raise ValueError( + f"Container {containers[decl]} of declaration {decl} does " + f"not contain all requested objects: {objects}." + ) + + def empty(self) -> bool: + return sum(map(len, self.requests.values())) == 0 + + def __iter__(self): + return self.requests.__iter__() + + def __getitem__(self, key: ContainerDeclaration) -> ObjectSet: + return self.requests[key] + + def __setitem__(self, key: ContainerDeclaration, value: ObjectSet) -> None: + if not isinstance(value, ObjectSet): + raise TypeError(f"Expected ObjectSet, got {type(value)}") + self.requests[key] = value + + def __delitem__(self, key: ContainerDeclaration) -> None: + if key in self.requests: + del self.requests[key] + else: + raise KeyError(f"Container {key} not found in requests.") + + def __len__(self) -> int: + return len(self.requests) + + def insert(self, container: ContainerDeclaration, object_list: ObjectSet) -> None: + """ + Insert a new request for the given container and object list. + This is similar to update, but it does not check if the container already exists. + """ + if not isinstance(object_list, ObjectSet): + raise TypeError(f"Expected ObjectSet, got {type(object_list)}") + self._requests_for(container).update(object_list) + + def __repr__(self): + return repr(dict(self.requests)) + + def __eq__(self, other) -> bool: + return self.requests == other.requests diff --git a/python/revng/pypeline/task/savepoint.py b/python/revng/pypeline/task/savepoint.py new file mode 100644 index 000000000..5853c6689 --- /dev/null +++ b/python/revng/pypeline/task/savepoint.py @@ -0,0 +1,114 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from typing import Sequence, final + +from revng.pypeline.container import ConfigurationId, ContainerDeclaration, ContainerSet +from revng.pypeline.storage.storage_provider import ContainerLocation, SavePointsRange +from revng.pypeline.storage.storage_provider import StorageProvider + +from .requests import Requests +from .task import TaskArgument, TaskArgumentAccess + + +class SavePoint: + """ + A SavePoint is a task that satisfies its requests by looking up the requested objects in the + container it owns. + """ + + __slots__: tuple = ("name", "arguments", "to_save") + + def __init__( + self, + name: str, + to_save: Sequence[ContainerDeclaration], + ): + self.name = f"SavePoint({name or self.__class__.__name__})" + self.arguments = [ + TaskArgument( + name=declaration.name, + container_type=declaration.container_type, + access=TaskArgumentAccess.READ_WRITE, + ) + for declaration in to_save + ] + self.to_save = to_save + + @final + def prerequisites_for( + self, + requests: Requests, + configuration_id: ConfigurationId, + storage_provider: StorageProvider, + savepoint_range: SavePointsRange, + ) -> Requests: + """ + This method has the same semantics as Pipe.prerequisites_for, but it has + additional arguments that the Pipe should not have. + """ + result = requests.clone() + + for decl, request in result.items(): + if decl not in self.to_save: + # This savepoint does not care about this container + continue + # We have a match, so we can check if the objects are present + location = ContainerLocation( + savepoint_id=savepoint_range.start, + container_id=decl.name, + configuration_id=configuration_id, + ) + found_objs = storage_provider.has(location, request) + # If so, we should REMOVE the found objects from the request + for found_obj in found_objs: + result[decl].remove(found_obj) + return result + + def run( + self, + containers: ContainerSet, + incoming: Requests, + outgoing: Requests, + configuration_id: ConfigurationId, + storage_provider: StorageProvider, + savepoint_range: SavePointsRange, + ): + """ + The savepoint will cache the incoming containers, and then fill the outgoing + containers with the cached data. + """ + + assert set(containers.keys()).issuperset( + set(incoming.keys()) | set(outgoing.keys()) + ), "SavePoint containers must be a subset of incoming and outgoing requests." + + # Cache the containers present for this configuration + for decl in self.to_save: + if decl not in incoming: + continue + if len(incoming[decl]) == 0: + continue + storage_provider.put( + ContainerLocation( + savepoint_id=savepoint_range.start, + container_id=decl.name, + configuration_id=configuration_id, + ), + containers[decl].serialize(incoming[decl]), + ) + + # Fill the containers from our cache + for decl in self.to_save: + if decl not in outgoing: + continue + # Empty requests do not need to be restored + if len(outgoing[decl]) == 0: + continue + location = ContainerLocation( + savepoint_id=savepoint_range.start, + container_id=decl.name, + configuration_id=configuration_id, + ) + containers[decl].deserialize(storage_provider.get(location, outgoing[decl])) diff --git a/python/revng/pypeline/task/task.py b/python/revng/pypeline/task/task.py new file mode 100644 index 000000000..3f547aa0a --- /dev/null +++ b/python/revng/pypeline/task/task.py @@ -0,0 +1,55 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from dataclasses import dataclass +from enum import Flag, auto +from typing import Annotated + +from revng.pypeline.container import ContainerDeclaration, ContainerID +from revng.pypeline.model import ModelPath +from revng.pypeline.object import ObjectID + +PipeObjectDependencies = Annotated[ + list[list[tuple[ObjectID, ModelPath]]], + """ + A list representing the dependencies between the an object (in a certain container) produced + by a Pipe. As the Pipe doesn't know the container names, it just returns + the index of the container in the Pipe's signature. And then it's + up to `PipelineNode` to remap the index to the container name. + """, +] + +ObjectDependencies = Annotated[ + list[tuple[ContainerID, ObjectID, ModelPath]], + """ + A list representing the dependencies between the an object (in a certain container) produced + by a certain task and the model. + """, +] + + +class TaskArgumentAccess(Flag): + READ = auto() + WRITE = auto() + READ_WRITE = READ | WRITE + + +@dataclass(slots=True, frozen=True) +class TaskArgument(ContainerDeclaration): + """ + A container argument for a task. + + It has a name, it has to be bounded to a ContainerDeclaration and we can specify whether the + container is just read, just written or both. + """ + + access: TaskArgumentAccess + # This is the description of the argument that will appear in the CLI + help_text: str = "" + + def to_container_decl(self) -> ContainerDeclaration: + """ + Convert this TaskArgument to a ContainerDeclaration. + """ + return ContainerDeclaration(self.name, self.container_type) diff --git a/python/revng/pypeline/utils/__init__.py b/python/revng/pypeline/utils/__init__.py new file mode 100644 index 000000000..79d95cc20 --- /dev/null +++ b/python/revng/pypeline/utils/__init__.py @@ -0,0 +1,3 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# diff --git a/python/revng/pypeline/utils/cabc.py b/python/revng/pypeline/utils/cabc.py new file mode 100644 index 000000000..316b30b76 --- /dev/null +++ b/python/revng/pypeline/utils/cabc.py @@ -0,0 +1,30 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +""" +CABC: Compatibility ABCs. + +This is needed to support proper type checking, but +avoid metaclasses which we should avoid on classes we intend of implementing in +C++. +""" + +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from abc import ABC, abstractmethod +else: + + class ABC: + pass + + def abstractmethod(func): + # Import it inside so users cannot import it by mistake + from functools import wraps + + @wraps(func) + def wrapper(*args, **kwargs): + raise NotImplementedError("Abstract method") + + return wrapper diff --git a/python/revng/pypeline/utils/default_dict_from_key.py b/python/revng/pypeline/utils/default_dict_from_key.py new file mode 100644 index 000000000..737afadef --- /dev/null +++ b/python/revng/pypeline/utils/default_dict_from_key.py @@ -0,0 +1,31 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +from collections.abc import Mapping +from typing import Callable, Dict, Generic, TypeVar + +K = TypeVar("K") +V = TypeVar("V") + + +class DefaultDictFromKey(Generic[K, V], Mapping[K, V]): + """ + A dictionary that, when queried with a key not currently present in the dictionary, + automatically creates it initializing it using a custom function of the key. + """ + + def __init__(self, factory: Callable[[K], V]): + self._dictionary: Dict[K, V] = {} + self._factory = factory + + def __getitem__(self, key: K) -> V: + if key not in self._dictionary: + self._dictionary[key] = self._factory(key) + return self._dictionary[key] + + def __iter__(self): + return self._dictionary.__iter__() + + def __len__(self): + return self._dictionary.__len__() diff --git a/python/revng/pypeline/utils/registry.py b/python/revng/pypeline/utils/registry.py new file mode 100644 index 000000000..a0d900dab --- /dev/null +++ b/python/revng/pypeline/utils/registry.py @@ -0,0 +1,61 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +# The singleton registry of subclasses for all classes decorated with `registry` +_REGISTRY_SINGLETON: dict[type, dict[str, type]] = {} + + +def get_registry[T](cls: type[T]) -> dict[str, type[T]]: + """ + Get the registry of subclasses for a given class which was decorated with + `registry`. + """ + if cls not in _REGISTRY_SINGLETON: + raise TypeError("`cls` is not a registered class") + # Clone to avoid mutation + return dict(_REGISTRY_SINGLETON[cls]) + + +def get_singleton[T](cls: type[T]) -> type[T]: + """ + Get the singleton subclass for a given class which was decorated with + `registry`. + """ + ty_registry = get_registry(cls) + if len(ty_registry) != 1: + raise ValueError( + f"Expected exactly one singleton for {cls}, but found " + f"{len(ty_registry)}: {list(ty_registry.keys())}" + ) + # Get the single value from the registry + return next(iter(ty_registry.values())) + + +def register_all_subclasses(cls: type, *, singleton: bool = False) -> None: + """ + Register all subclasses of the given class `cls` in the global registry. + """ + if cls not in _REGISTRY_SINGLETON: + _REGISTRY_SINGLETON[cls] = {} + + def find_leafs(c: type) -> list[type]: + """ + Recursively find all leaf subclasses of a given class `c`. + """ + subclasses = c.__subclasses__() + if not subclasses: + return [c] + leafs = [] + for subclass in subclasses: + leafs.extend(find_leafs(subclass)) + return leafs + + for leaf in find_leafs(cls): + _REGISTRY_SINGLETON[cls][leaf.__name__] = leaf + + if singleton and len(_REGISTRY_SINGLETON[cls]) > 1: + raise ValueError( + f"Expected exactly one singleton for {cls}, but found " + f"{len(_REGISTRY_SINGLETON[cls])}: {list(_REGISTRY_SINGLETON[cls].keys())}" + ) diff --git a/share/revng/rcc-config.yml b/share/revng/rcc-config.yml index 9951df1bf..89d1cd809 100644 --- a/share/revng/rcc-config.yml +++ b/share/revng/rcc-config.yml @@ -68,6 +68,10 @@ read_passes: # inference on untyped functions - --no-check-untyped-defs - --disable-error-code=annotation-unchecked + # This is needed so mypy doesn't complain about having a wheel called + # `revng` and having a file `revng.py` which implicitly defines a + # duplicated module called `revng`, but it works. + - --exclude="python/scripts/revng" - type: LicenseCheckPass ignore_suffixes: [".txt", ".md", ".rst", ".dot"] ignore_stems: ["LICENSE"] diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 990389715..b71c56ee2 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -4,4 +4,5 @@ add_subdirectory(abi) add_subdirectory(pipeline) +add_subdirectory(pypeline) add_subdirectory(unit) diff --git a/tests/pypeline/CMakeLists.txt b/tests/pypeline/CMakeLists.txt new file mode 100644 index 000000000..6c050c01c --- /dev/null +++ b/tests/pypeline/CMakeLists.txt @@ -0,0 +1,21 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +revng_add_test( + NAME + test_pypeline + COMMAND + "${Python_EXECUTABLE}" + -m + pytest + -p + no:cacheprovider + -v + --rootdir + "${CMAKE_CURRENT_SOURCE_DIR}" + "${CMAKE_CURRENT_SOURCE_DIR}/test_all.py") +set_tests_properties( + test_pypeline + PROPERTIES LABELS "pypeline" ENVIRONMENT + "PYTHONPATH=${CMAKE_BINARY_DIR}/${PYTHON_INSTALL_PATH}") diff --git a/tests/pypeline/model.yml b/tests/pypeline/model.yml new file mode 100644 index 000000000..1766d7d66 --- /dev/null +++ b/tests/pypeline/model.yml @@ -0,0 +1,11 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +/one: 1 +/two: 2 +/three: 3 +/all/one: 1 +/all/two: 2 +/all/three: 3 +/test/test: test diff --git a/tests/pypeline/pipeline.yml b/tests/pypeline/pipeline.yml new file mode 100644 index 000000000..ab12dc213 --- /dev/null +++ b/tests/pypeline/pipeline.yml @@ -0,0 +1,61 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# + +containers: + - name: child_source + type: ChildDictContainer + + - name: child_destination + type: ChildDictContainer + + - name: root_source + type: RootDictContainer + + - name: root_destination + type: RootDictContainer + +branches: + root: + tasks: + - pipe: GeneratorPipe + arguments: [child_source] + configuration: "" + analyses: + - name: init_analysis + analysis: AddStuffAnalysis + containers: [child_source] + - pipe: InPlacePipe + arguments: [child_source] + - pipe: SameKindPipe + arguments: [child_source, child_destination] + - savepoint: + name: "after_same_kind" + containers: [child_source, child_destination] + + second_branch: + from: root + tasks: + - pipe: ToLowerKindPipe + arguments: [child_destination, root_source] + artifacts: + - name: RootArtifact + container: root_source + analyses: + - name: blackhole + analysis: PurgeAllAnalysis + containers: [root_source] + - pipe: ToHigherKindPipe + arguments: [root_source, child_destination] + + third_branch: + from: root + tasks: + - pipe: InPlacePipe + arguments: [child_destination] + artifacts: + - name: ChildArtifact + container: child_destination + analyses: + - analysis: NullAnalysis + containers: [child_destination] diff --git a/tests/pypeline/simple_pipeline.py b/tests/pypeline/simple_pipeline.py new file mode 100644 index 000000000..ddd25e0d2 --- /dev/null +++ b/tests/pypeline/simple_pipeline.py @@ -0,0 +1,544 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# +# pylint: disable=too-many-positional-arguments + +from __future__ import annotations + +from abc import ABC, ABCMeta +from enum import Enum, EnumMeta, auto, unique +from typing import Any, Dict, Iterable, Mapping, Optional, TypeVar, Union, cast + +import yaml + +from revng.pypeline.analysis import Analysis +from revng.pypeline.container import Configuration, Container +from revng.pypeline.model import Model, ModelPath, ModelPathSet, ReadOnlyModel +from revng.pypeline.object import Kind, ObjectID, ObjectSet +from revng.pypeline.task.pipe import Pipe +from revng.pypeline.task.task import PipeObjectDependencies, TaskArgument, TaskArgumentAccess + +Value = Union[str, int] +T = TypeVar("T") + + +def mandatory(arg: Optional[T]) -> T: + assert arg is not None + return arg + + +def only(value: list[Any]) -> Any: + assert len(value) == 1 + return value[0] + + +class KindEnumMeta(ABCMeta, EnumMeta): + """A metaclass that combines ABCMeta and EnumMeta to allow + abstract methods in an Enum.""" + + +@unique +class MyKind(Kind, Enum, metaclass=KindEnumMeta): + ROOT = cast(Kind, auto()) + CHILD = cast(Kind, auto()) + GRANDCHILD = cast(Kind, auto()) + CHILD2 = cast(Kind, auto()) + + @classmethod + def kinds(cls) -> list[Kind]: + return list(cast(Iterable[Kind], cls)) + + def parent(self) -> Kind | None: + # Small perf optimization: this could be moved outside + hierarchy: dict[Kind, Kind | None] = { + self.ROOT: None, + self.CHILD: self.ROOT, + self.GRANDCHILD: self.CHILD, + self.CHILD2: self.ROOT, + } + return hierarchy[self] + + def __hash__(self) -> int: + return hash(self.name) + + def serialize(self) -> str: + return self.name + + @classmethod + def deserialize(cls, obj: str) -> Kind: + return MyKind[obj] + + +class MyObjectID(ObjectID): + """An object ID that is a sequence of strings.""" + + # Make ObjectIDs immutable + def __init__(self, kind: Kind, *components: str): + # A simple check to maintain structural integrity. + if kind.rank() != len(components): + raise ValueError("Number of components must match the kind's rank.") + assert isinstance(kind, MyKind), f"Expected kind to be a MyKind, got: {kind!r}" + self._kind: MyKind = kind + self._components = tuple(components) # Store as tuple to ensure immutability + + def kind(self) -> Kind: + return self._kind + + @classmethod + def root(cls) -> ObjectID: + return MyObjectID(MyKind.ROOT) + + def parent(self) -> Optional[ObjectID]: + parent_kind = self._kind.parent() + if parent_kind is None: + return None + return MyObjectID(parent_kind, *self._components[:-1]) + + # since we have singleton items, the comparisons and hashes + # can just be the object ptr + def __eq__(self, other) -> bool: + return hash(self) == hash(other) + + def __hash__(self) -> int: + return hash((self._kind, self._components)) + + def serialize(self) -> str: + return f"/{self._kind.name}/{'/'.join(self._components)}" + + @classmethod + def deserialize(cls, obj: str) -> ObjectID: + """Deserialize an object id of this class""" + kind, *components = obj.strip("/").split("/") + return cls(MyKind.deserialize(kind), *components) + + +class DictContainer(Container, ABC): + def __init__(self): + self._object_list: ObjectSet = ObjectSet(self.kind) + + def clone(self, objects: Optional[ObjectSet]) -> DictContainer: + # DictContainer is an abstract class that doesn't have a specific kind, + # So to create a new instance, we get the current "subclass" + # and instantiate it. + result = self.__class__() + if objects is not None: + assert objects.issubset(self._object_list) + result._object_list = objects # pylint: disable=protected-access + else: + result._object_list = self._object_list # pylint: disable=protected-access + return result + + def merge(self, other: Container) -> None: + # pylint: disable=protected-access + self._object_list.update(cast(DictContainer, other)._object_list) + + def objects(self) -> ObjectSet: + return ObjectSet(self._object_list.kind, set(self._object_list.objects)) + + def erase_objects(self, objects: ObjectSet): + self._object_list -= objects + + def clear(self): + self._object_list = ObjectSet(self.kind) + + def verify(self) -> bool: + return True + + def add_object(self, new_object: ObjectID): + self._object_list.add(new_object) + + def __repr__(self): + return f"DictContainer({self._object_list!r})" + + def __str__(self): + return repr(self) + + @classmethod + def mime_type(cls) -> str: + return "text" + + def deserialize(self, data: Mapping[ObjectID, bytes]) -> None: + for oid, _ in data.items(): + self.add_object(oid) + + def serialize(self, objects: Optional[ObjectSet] = None) -> dict[ObjectID, bytes]: + if objects is None: + return dict.fromkeys(self._object_list, b"") + return dict.fromkeys(objects.objects, b"") + + +class RootDictContainer(DictContainer): + kind = MyKind.ROOT + + +class ChildDictContainer(DictContainer): + kind = MyKind.CHILD + + +class DictModel(Model): + def __init__(self): + self._data: Dict[ModelPath, Value] = {} + + def __contains__(self, path: ModelPath) -> bool: + return path in self._data + + def __getitem__(self, path: ModelPath) -> Value: + return self._data[path] + + def __setitem__(self, path: ModelPath, new_value: Value): + self._data[path] = new_value + + def __delitem__(self, path: ModelPath): + if path in self._data: + del self._data[path] + else: + raise KeyError(f"ModelPath {path} not found in the model.") + + def __eq__(self, other: object) -> bool: + if not isinstance(other, DictModel): + return False + if set(self._data.keys()) != set(other._data.keys()): + return False + return all(self._data[key] == other._data[key] for key in self._data) + + def __len__(self) -> int: + return len(self._data) + + def items(self) -> list[tuple[ModelPath, Value]]: + return list(self._data.items()) + + def keys(self) -> list[ModelPath]: + return list(self._data.keys()) + + def values(self) -> list[Value]: + return list(self._data.values()) + + def diff(self, other: DictModel) -> ModelPathSet: + diff: ModelPathSet = set() + for key, value in self.items(): + if key not in other or other[key] != value: + diff.add(key) + return diff + + def clone(self) -> DictModel: + result = DictModel() + result._data = dict(self._data) # pylint: disable=protected-access + return result + + def children(self, obj: ObjectID, kind: Kind) -> ObjectSet: + if obj.kind() == MyKind.ROOT: + if kind == MyKind.CHILD: + return ObjectSet( + MyKind.CHILD, + { + MyObjectID(MyKind.CHILD, "one"), + MyObjectID(MyKind.CHILD, "two"), + MyObjectID(MyKind.CHILD, "three"), + }, + ) + elif kind == MyKind.ROOT: + return ObjectSet(MyKind.ROOT, {MyObjectID(MyKind.ROOT)}) + + raise NotImplementedError() + + @classmethod + def is_text(cls) -> bool: + # We serialize already as json + return True + + def serialize(self): + return yaml.safe_dump(self._data).encode() + + def deserialize(self, data: bytes): + self._data = yaml.safe_load(data) + + def __repr__(self): + return f"DictModel({self._data!r})" + + def __str__(self): + return repr(self) + + +class InPlacePipe(Pipe): + """Modifies the input container in place.""" + + @classmethod + def signature(cls) -> tuple[TaskArgument, ...]: + return ( + TaskArgument( + "arg", + ChildDictContainer, + TaskArgumentAccess.READ_WRITE, + help_text="the input container the pipe will modify", + ), + ) + + def run( + self, + model: ReadOnlyModel, + containers: list[Container], + incoming: list[ObjectSet], + outgoing: list[ObjectSet], + configuration: Configuration, + ) -> PipeObjectDependencies: + # Nothing to do + return [[]] + + +class SameKindPipe(Pipe): + @classmethod + def signature(cls) -> tuple[TaskArgument, ...]: + return ( + TaskArgument( + "source", + ChildDictContainer, + TaskArgumentAccess.READ, + help_text="the source container", + ), + TaskArgument( + "destination", + ChildDictContainer, + TaskArgumentAccess.WRITE, + help_text="the destination container", + ), + ) + + def run( + self, + model: ReadOnlyModel, + containers: list[Container], + incoming: list[ObjectSet], + outgoing: list[ObjectSet], + configuration: Configuration, + ) -> PipeObjectDependencies: + input_container: ChildDictContainer = cast(ChildDictContainer, containers[0]) + output_container: ChildDictContainer = cast(ChildDictContainer, containers[1]) + + for obj in input_container.objects().objects: + output_container.add_object(obj) + + return [[], []] + + +class ToHigherKindPipe(Pipe): + """Take the root object from the input container and adds all + its children to the output container.""" + + @classmethod + def signature(cls) -> tuple[TaskArgument, ...]: + return ( + TaskArgument( + "source", + RootDictContainer, + TaskArgumentAccess.READ, + help_text="the source container", + ), + TaskArgument( + "destination", + ChildDictContainer, + TaskArgumentAccess.WRITE, + help_text="the destination container", + ), + ) + + def run( + self, + model: ReadOnlyModel, + containers: list[Container], + incoming: list[ObjectSet], + outgoing: list[ObjectSet], + configuration: Configuration, + ) -> PipeObjectDependencies: + input_container: RootDictContainer = cast(RootDictContainer, containers[0]) + input_kind = RootDictContainer.kind + + # Ensure we have the root object in input + assert input_container.objects() == ObjectSet(input_kind, {MyObjectID.root()}), ( + f"Expected input container to contain only the root object, got: " + f"{input_container.objects()}" + ) + + # Add all the children of the root object in output + output_container: ChildDictContainer = cast(ChildDictContainer, containers[1]) + output_kind = output_container.kind + root_object = MyObjectID.root() + for obj in model.children(root_object, output_kind).objects: + output_container.add_object(obj) + + return [[], []] + + +class ToLowerKindPipe(Pipe): + @classmethod + def signature(cls) -> tuple[TaskArgument, ...]: + return ( + TaskArgument( + "source", + ChildDictContainer, + TaskArgumentAccess.READ, + help_text="the source container", + ), + TaskArgument( + "destination", + RootDictContainer, + TaskArgumentAccess.WRITE, + help_text="the destination container", + ), + ) + + def run( + self, + model: ReadOnlyModel, + containers: list[Container], + incoming: list[ObjectSet], + outgoing: list[ObjectSet], + configuration: Configuration, + ) -> PipeObjectDependencies: + input_container: ChildDictContainer = cast(ChildDictContainer, containers[0]) + input_kind = input_container.kind + root_object = MyObjectID.root() + + # Ensure we have all the object we need in input + assert input_container.objects() == model.children(root_object, input_kind) + + output_container: RootDictContainer = cast(RootDictContainer, containers[1]) + + # Add to the output the root object + output_container.add_object(MyObjectID.root()) + + return [[], []] + + +class GeneratorPipe(Pipe): + @classmethod + def signature(cls) -> tuple[TaskArgument, ...]: + return ( + TaskArgument( + "arg", + ChildDictContainer, + TaskArgumentAccess.WRITE, + help_text="the output container", + ), + ) + + def __init__( + self, + static_configuration: str = "", + name: str | None = None, + ): + super().__init__( + name=name, + static_configuration=static_configuration, + ) + + def run( + self, + model: ReadOnlyModel, + containers: list[Container], + incoming: list[ObjectSet], + outgoing: list[ObjectSet], + configuration: Configuration, + ) -> PipeObjectDependencies: + dependencies: PipeObjectDependencies = [] + model = model.downcast() + assert isinstance(model, DictModel), f"Model must be a DictModel got: {model!r}" + container = cast(ChildDictContainer, containers[0]) + + for objects in outgoing: + for obj in objects: + container.add_object(obj) + dependencies.append((obj, "/one")) + return [dependencies] + + +class NullAnalysis(Analysis): + """An analysis that does nothing and returns an empty list of invalidations.""" + + @classmethod + def signature(cls) -> tuple[type[Container], ...]: + return (ChildDictContainer,) + + def run( + self, + model: Model, + containers: list[Container], + incoming: list[ObjectSet], + configuration: str, + ): + # This analysis does nothing + pass + + +class PurgeOneAnalysis(Analysis): + """An analysis that invalidates everything.""" + + @classmethod + def signature(cls) -> tuple[type[Container], ...]: + return (ChildDictContainer,) + + def __init__(self, name: str): + super().__init__(name) + self.what_to_purge: list[ModelPath] = [ + "/one", + "/test/test", + ] + + def run( + self, + model: Model, + containers: list[Container], + incoming: list[ObjectSet], + configuration: str, + ): + assert isinstance(model, DictModel) + for purge_path in self.what_to_purge: + if purge_path in model: + del model[purge_path] + + +class PurgeAllAnalysis(Analysis): + """An analysis that invalidates everything.""" + + @classmethod + def signature(cls) -> tuple[type[Container], ...]: + return (ChildDictContainer,) + + def run( + self, + model: Model, + containers: list[Container], + incoming: list[ObjectSet], + configuration: str, + ): + assert isinstance(model, DictModel) + keys = model.keys() + for key in keys: + del model[key] + + +class AddStuffAnalysis(Analysis): + """An analysis that invalidates everything.""" + + @classmethod + def signature(cls) -> tuple[type[Container], ...]: + return (ChildDictContainer,) + + def __init__(self, name: str): + super().__init__(name) + self.what_to_add: list[ModelPath] = [ + "/one", + "/test/test", + "/test/hello", + ] + + def run( + self, + model: Model, + containers: list[Container], + incoming: list[ObjectSet], + configuration: str, + ): + assert isinstance(model, DictModel) + for add_path in self.what_to_add: + if add_path not in model: + model[add_path] = f"wooo {add_path}" diff --git a/tests/pypeline/test_all.py b/tests/pypeline/test_all.py new file mode 100644 index 000000000..f38d56da5 --- /dev/null +++ b/tests/pypeline/test_all.py @@ -0,0 +1,648 @@ +# +# This file is distributed under the MIT License. See LICENSE.md for details. +# +# pylint: disable=redefined-outer-name + +from __future__ import annotations + +import os +from tempfile import NamedTemporaryFile +from typing import Optional, TypeVar, Union + +import pytest +from simple_pipeline import ChildDictContainer, DictModel, GeneratorPipe, InPlacePipe, MyKind +from simple_pipeline import MyObjectID, NullAnalysis, PurgeAllAnalysis, PurgeOneAnalysis +from simple_pipeline import RootDictContainer, SameKindPipe, ToHigherKindPipe, ToLowerKindPipe + +from revng.pypeline import initialize_pypeline +from revng.pypeline.analysis import AnalysisBinding +from revng.pypeline.container import ContainerDeclaration +from revng.pypeline.model import ReadOnlyModel +from revng.pypeline.object import Kind, ObjectSet +from revng.pypeline.pipeline import Artifact, Pipeline +from revng.pypeline.pipeline_node import PipelineConfiguration, PipelineNode +from revng.pypeline.pipeline_parser import load_pipeline_yaml_file +from revng.pypeline.storage.memory import InMemoryStorageProvider +from revng.pypeline.storage.sqlite3 import SQlite3StorageProvider +from revng.pypeline.storage.storage_provider import ContainerLocation, SavePointsRange +from revng.pypeline.storage.storage_provider import StorageProvider +from revng.pypeline.task.pipe import Pipe +from revng.pypeline.task.requests import Requests +from revng.pypeline.task.savepoint import SavePoint + +# Fill the registries +initialize_pypeline() + +Value = Union[str, int] + + +T = TypeVar("T") + + +def mandatory(arg: Optional[T]) -> T: + assert arg is not None + return arg + + +@pytest.fixture +def model(): + return DictModel() + + +@pytest.fixture(params=["memory", "sqlite3"]) +def storage_provider(request): + storage_provider: StorageProvider + if request.param == "memory": + storage_provider = InMemoryStorageProvider() + yield storage_provider + elif request.param == "sqlite3": + with NamedTemporaryFile() as f: + storage_provider = SQlite3StorageProvider(":memory:", f.name) + yield storage_provider + else: + raise ValueError() + + +def test_kind(): + # Test rank + assert MyKind.ROOT.rank() == 0 + assert MyKind.CHILD.rank() == 1 + assert MyKind.GRANDCHILD.rank() == 2 + assert MyKind.CHILD2.rank() == 1 + assert MyKind.root() == MyKind.ROOT + assert MyKind.CHILD.parent() == MyKind.ROOT + assert MyKind.GRANDCHILD.parent() == MyKind.CHILD + assert MyKind.CHILD2.parent() == MyKind.ROOT + assert MyKind.kinds() == [MyKind.ROOT, MyKind.CHILD, MyKind.GRANDCHILD, MyKind.CHILD2] + + # Test relation + assert MyKind.ROOT.relation(MyKind.ROOT)[0] == Kind.Relation.SAME + + assert MyKind.CHILD.relation(MyKind.CHILD2)[0] == Kind.Relation.UNRELATED + + assert MyKind.ROOT.relation(MyKind.CHILD) == ( + Kind.Relation.ANCESTOR, + [MyKind.ROOT, MyKind.CHILD], + ) + assert MyKind.CHILD.relation(MyKind.ROOT) == ( + Kind.Relation.DESCENDANT, + [MyKind.CHILD, MyKind.ROOT], + ) + + assert MyKind.ROOT.relation(MyKind.GRANDCHILD) == ( + Kind.Relation.ANCESTOR, + [MyKind.ROOT, MyKind.CHILD, MyKind.GRANDCHILD], + ) + assert MyKind.GRANDCHILD.relation(MyKind.ROOT) == ( + Kind.Relation.DESCENDANT, + [MyKind.GRANDCHILD, MyKind.CHILD, MyKind.ROOT], + ) + + +def test_pipe_prerequisites_for(model) -> None: + # These are special declarations and have to exactly match the args of the + # pipes being tested. + root = ObjectSet(MyKind.ROOT, {MyObjectID.root()}) + one_two = ObjectSet( + MyKind.CHILD, {MyObjectID(MyKind.CHILD, "one"), MyObjectID(MyKind.CHILD, "two")} + ) + one_two_three = ObjectSet( + MyKind.CHILD, + { + MyObjectID(MyKind.CHILD, "one"), + MyObjectID(MyKind.CHILD, "two"), + MyObjectID(MyKind.CHILD, "three"), + }, + ) + empty_child_set = ObjectSet(MyKind.CHILD, set()) + empty_root_set = ObjectSet(MyKind.ROOT, set()) + + pipe: Pipe + + pipe = InPlacePipe() + result = pipe.prerequisites_for(model, [one_two]) + assert result == [one_two] + + pipe = SameKindPipe() + result = pipe.prerequisites_for(model, [empty_child_set, one_two]) + assert result == [one_two, empty_child_set] + + pipe = ToLowerKindPipe() + result = pipe.prerequisites_for(model, [empty_child_set, root]) + assert result == [one_two_three, empty_root_set] + + pipe = ToHigherKindPipe() + result = pipe.prerequisites_for(model, [empty_root_set, one_two_three]) + assert result == [root, empty_child_set] + + +def test_savepoint_prerequisites_for(storage_provider, model) -> None: + storage_provider.set_model(model.serialize()) + child: ContainerDeclaration = ContainerDeclaration("child", ChildDictContainer) + save_point = SavePoint("save", [child]) + + container = ChildDictContainer() + container.add_object(MyObjectID(MyKind.CHILD, "one")) + requests = Requests({child: ObjectSet(MyKind.CHILD, {MyObjectID(MyKind.CHILD, "one")})}) + savepoint_range = SavePointsRange(start=0, end=0) + configuration_id = "132124ujh12jk4jk124kj" + + # The storage is empty so the savepoint should be transparent + result = save_point.prerequisites_for( + requests=requests, + configuration_id=configuration_id, + storage_provider=storage_provider, + savepoint_range=savepoint_range, + ) + assert result == requests + + # Force the savepoint to store the object + save_point.run( + containers={child: container}, + incoming=requests, + outgoing=requests, + configuration_id=configuration_id, + storage_provider=storage_provider, + savepoint_range=savepoint_range, + ) + + result = save_point.prerequisites_for( + requests=requests, + configuration_id=configuration_id, + storage_provider=storage_provider, + savepoint_range=savepoint_range, + ) + assert result == Requests({child: ObjectSet(MyKind.CHILD, set())}) + + result = save_point.prerequisites_for( + requests=Requests( + { + child: ObjectSet( + MyKind.CHILD, {MyObjectID(MyKind.CHILD, "one"), MyObjectID(MyKind.CHILD, "two")} + ) + } + ), + configuration_id=configuration_id, + storage_provider=storage_provider, + savepoint_range=savepoint_range, + ) + expected = Requests({child: ObjectSet(MyKind.CHILD, {MyObjectID(MyKind.CHILD, "two")})}) + assert result == expected + + +def test_pipeline_inplace(model, storage_provider): + child_cont: ContainerDeclaration = ContainerDeclaration("arg", ChildDictContainer) + declarations = [child_cont] + pipeline_configuration: PipelineConfiguration = {} + storage_provider.set_model(model.serialize()) + + one = ObjectSet(MyKind.CHILD, {MyObjectID(MyKind.CHILD, "one")}) + one_two = ObjectSet( + MyKind.CHILD, {MyObjectID(MyKind.CHILD, "one"), MyObjectID(MyKind.CHILD, "two")} + ) + + # Create the pipeline + begin_node = PipelineNode(SavePoint("begin", to_save=declarations)) + inplace_node = PipelineNode(InPlacePipe(), bindings=[child_cont]) + end_node = PipelineNode(SavePoint("end", to_save=declarations)) + begin_node.add_successor(inplace_node).add_successor(end_node) + + pipeline: Pipeline = Pipeline(set(declarations), begin_node) + assert begin_node.savepoint_range == SavePointsRange(start=1, end=2) + assert inplace_node.savepoint_range == SavePointsRange(start=1, end=2) + assert end_node.savepoint_range == SavePointsRange(start=2, end=2) + + # Force the savepoint to store the objects + container = ChildDictContainer() + container.add_object(MyObjectID(MyKind.CHILD, "one")) + container.add_object(MyObjectID(MyKind.CHILD, "two")) + begin_configuration_id = begin_node.configuration_id(pipeline_configuration) + begin_node.run( + model=ReadOnlyModel(model), + containers={child_cont: container}, + incoming=Requests({child_cont: one_two}), + outgoing=Requests({child_cont: one_two}), + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + + assert set( + storage_provider.has( + location=ContainerLocation( + savepoint_id=1, + container_id="arg", + configuration_id=begin_configuration_id, + ), + keys=one_two, + ) + ) == set(one_two) + + containers = pipeline.schedule( + model=ReadOnlyModel(model), + target_node=end_node, + requests=Requests({child_cont: one}), + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ).run( + model=ReadOnlyModel(model), + storage_provider=storage_provider, + ) + assert containers[child_cont].objects() == one + + end_configuration_id = end_node.configuration_id(pipeline_configuration) + assert list( + storage_provider.has( + location=ContainerLocation( + savepoint_id=2, + container_id="arg", + configuration_id=end_configuration_id, + ), + keys=one, + ) + ) == list(one) + + +def test_pipeline_up_down(model, storage_provider): + root1: ContainerDeclaration = ContainerDeclaration("root_source", RootDictContainer) + root2: ContainerDeclaration = ContainerDeclaration("root_destination", RootDictContainer) + child1: ContainerDeclaration = ContainerDeclaration("child_destination", ChildDictContainer) + child2: ContainerDeclaration = ContainerDeclaration("child_source", ChildDictContainer) + + declarations = [root1, root2, child1, child2] + + pipeline_configuration: PipelineConfiguration = {} + storage_provider.set_model(model.serialize()) + + begin_node = PipelineNode(SavePoint("begin", declarations)) + up_node = PipelineNode(ToHigherKindPipe(), bindings=[root1, child1]) + same_node = PipelineNode(SameKindPipe(), bindings=[child1, child2]) + down_node = PipelineNode(ToLowerKindPipe(), bindings=[child2, root2]) + end_node = PipelineNode(SavePoint("end", declarations)) + begin_node.add_successor(up_node).add_successor(same_node).add_successor( + down_node + ).add_successor(end_node) + + pipeline: Pipeline = Pipeline(set(declarations), begin_node) + assert begin_node.savepoint_range == SavePointsRange(start=1, end=2) + assert up_node.savepoint_range == SavePointsRange(start=1, end=2) + assert same_node.savepoint_range == SavePointsRange(start=1, end=2) + assert down_node.savepoint_range == SavePointsRange(start=1, end=2) + assert end_node.savepoint_range == SavePointsRange(start=2, end=2) + + root_obj = ObjectSet(MyKind.ROOT, {MyObjectID.root()}) + + # Force the savepoint to store the objects + container = RootDictContainer() + container.add_object(MyObjectID.root()) + requests = Requests({root1: root_obj}) + begin_node.run( + model=ReadOnlyModel(model), + containers={root1: container}, + incoming=requests, + outgoing=requests, + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + + containers = pipeline.schedule( + model=ReadOnlyModel(model), + target_node=end_node, + requests=Requests({root2: root_obj}), + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ).run( + model=ReadOnlyModel(model), + storage_provider=storage_provider, + ) + assert containers[root2].objects() == root_obj + + end_configuration_id = end_node.configuration_id(pipeline_configuration) + assert set( + storage_provider.has( + location=ContainerLocation( + savepoint_id=2, + container_id="root_destination", + configuration_id=end_configuration_id, + ), + keys=[MyObjectID.root()], + ) + ) == {MyObjectID.root()} + + +def test_artifact(model, storage_provider): + child1: ContainerDeclaration = ContainerDeclaration("source", ChildDictContainer) + child2: ContainerDeclaration = ContainerDeclaration("destination", ChildDictContainer) + declarations = [child1, child2] + pipeline_configuration: PipelineConfiguration = {} + one = ObjectSet.from_list([MyObjectID(MyKind.CHILD, "one")]) + storage_provider.set_model(model.serialize()) + + begin_node = PipelineNode(SavePoint("begin", declarations)) + same_node = PipelineNode(SameKindPipe(), bindings=[child1, child2]) + begin_node.add_successor(same_node) + + artifact = Artifact("artifact", same_node, child2) + + pipeline: Pipeline = Pipeline(set(declarations), begin_node) + assert begin_node.savepoint_range == SavePointsRange(start=1, end=1) + assert same_node.savepoint_range == SavePointsRange(start=1, end=1) + + # Force the savepoint to store the objects + container = ChildDictContainer() + container.add_object(MyObjectID(MyKind.CHILD, "one")) + requests = Requests({child1: one}) + begin_node.run( + model=ReadOnlyModel(model), + containers={child1: container}, + incoming=requests, + outgoing=requests, + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + + res = pipeline.get_artifact( + model=ReadOnlyModel(model), + artifact=artifact, + requests=one, + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + + assert res.objects() == one + + +def test_invalidation(model, storage_provider): + # TODO: simple test: one pipe, one save point, multiple objects, invalidate one + child: ContainerDeclaration = ContainerDeclaration("arg", ChildDictContainer) + declarations = [child] + pipeline_configuration: PipelineConfiguration = {} + storage_provider.set_model(model.serialize()) + + object_one = MyObjectID(MyKind.CHILD, "one") + expected_output = ObjectSet(MyKind.CHILD, {object_one}) + + pipe = PipelineNode(GeneratorPipe(), bindings=[child]) + savepoint = PipelineNode(SavePoint("end", declarations)) + pipe.add_successor(savepoint) + + pipeline: Pipeline = Pipeline(set(declarations), pipe) + assert pipe.savepoint_range == SavePointsRange(start=0, end=1) + assert savepoint.savepoint_range == SavePointsRange(start=1, end=1) + + containers = pipeline.schedule( + model=ReadOnlyModel(model), + target_node=savepoint, + requests=Requests({child: expected_output}), + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ).run( + model=ReadOnlyModel(model), + storage_provider=storage_provider, + ) + assert containers[child].objects() == expected_output + + assert set( + storage_provider.has( + location=ContainerLocation( + savepoint_id=1, + container_id="arg", + configuration_id=savepoint.configuration_id(pipeline_configuration), + ), + keys=expected_output, + ) + ) == set(expected_output) + + storage_provider.invalidate({"/two"}) + + assert set( + storage_provider.has( + location=ContainerLocation( + savepoint_id=1, + container_id="arg", + configuration_id=savepoint.configuration_id(pipeline_configuration), + ), + keys=expected_output, + ) + ) == set(expected_output) + + # Change something we depend on + storage_provider.invalidate({"/one"}) + + assert not list( + storage_provider.has( + location=ContainerLocation( + savepoint_id=1, + container_id="arg", + configuration_id=savepoint.configuration_id(pipeline_configuration), + ), + keys=expected_output, + ) + ) + + +def test_analysis(model, storage_provider): + # TODO: simple test: one pipe, one save point, multiple objects, invalidate one + child: ContainerDeclaration = ContainerDeclaration("arg", ChildDictContainer) + declarations = [child] + pipeline_configuration: PipelineConfiguration = {} + + model["/test/test"] = "test" + model["/test/test2"] = "test2" + + storage_provider.set_model(model.serialize()) + + object_one = MyObjectID(MyKind.CHILD, "one") + expected_output = ObjectSet(MyKind.CHILD, {object_one}) + + pipe = PipelineNode( + GeneratorPipe(), + bindings=[child], + ) + savepoint = PipelineNode(SavePoint("end", declarations)) + pipe.add_successor(savepoint) + + pipeline: Pipeline = Pipeline( + set(declarations), + pipe, + analyses={ + AnalysisBinding( + NullAnalysis(name="null_analysis"), + (child,), + savepoint, + ), + AnalysisBinding( + PurgeAllAnalysis(name="purge_all_analysis"), + (child,), + savepoint, + ), + AnalysisBinding( + PurgeOneAnalysis(name="purge_one_analysis"), + (child,), + savepoint, + ), + }, + ) + assert pipe.savepoint_range == SavePointsRange(start=0, end=1) + assert savepoint.savepoint_range == SavePointsRange(start=1, end=1) + + orig_model = model.clone() + new_model = pipeline.run_analysis( + model=ReadOnlyModel(model), + analysis_name="null_analysis", + requests=Requests({child: expected_output}), + analysis_configuration="", + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + + assert model == orig_model, "Model should not be modified by the analysis" + assert new_model == orig_model, "This analysis doesn't invalidate anything" + + new_model = pipeline.run_analysis( + model=ReadOnlyModel(model), + analysis_name="purge_all_analysis", + requests=Requests({child: expected_output}), + analysis_configuration="", + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + + assert model == orig_model, "Model should not be modified by the analysis" + assert new_model == DictModel(), "This analysis invalidates everything" + + new_model = pipeline.run_analysis( + model=ReadOnlyModel(model), + analysis_name="purge_one_analysis", + requests=Requests({child: expected_output}), + analysis_configuration="", + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + + assert model == orig_model, "Model should not be modified by the analysis" + expected_model = DictModel() + expected_model["/test/test2"] = "test2" + assert new_model == expected_model, "This analysis invalidates everything" + + +def test_pipeline(storage_provider): + """Load the schema and validate the pipeline.yml file against it.""" + root = os.path.dirname(os.path.abspath(__file__)) + pipeline = load_pipeline_yaml_file(os.path.join(root, "pipeline.yml")) + pipeline_configuration: PipelineConfiguration = {} + + model = DictModel() + with open(os.path.join(root, "model.yml"), "rb") as model_file: + model.deserialize(model_file.read()) + + res = pipeline.get_artifact( + model=ReadOnlyModel(model), + artifact=pipeline.artifacts["ChildArtifact"], + requests=ObjectSet(MyKind.CHILD, {MyObjectID(MyKind.CHILD, "one")}), + pipeline_configuration={}, + storage_provider=storage_provider, + ) + assert res.objects() == ObjectSet(MyKind.CHILD, {MyObjectID(MyKind.CHILD, "one")}) + + res = pipeline.get_artifact( + model=ReadOnlyModel(model), + artifact=pipeline.artifacts["RootArtifact"], + requests=ObjectSet(MyKind.ROOT, {MyObjectID.root()}), + pipeline_configuration={}, + storage_provider=storage_provider, + ) + assert res.objects() == ObjectSet(MyKind.ROOT, {MyObjectID.root()}) + + new_model = pipeline.run_analysis( + model=ReadOnlyModel(model), + analysis_name="NullAnalysis", + requests=Requests( + { + ContainerDeclaration( + name="child_destination", + container_type=ChildDictContainer, + ): ObjectSet(MyKind.CHILD, {MyObjectID(MyKind.CHILD, "one")}) + } + ), + analysis_configuration="", + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + assert isinstance(new_model, DictModel), "The analysis should return the same model type" + assert new_model == model, "NullAnalysis should not change the model" + + new_model = pipeline.run_analysis( + model=ReadOnlyModel(model), + # An alias of PurgeAllAnalysis + analysis_name="blackhole", + requests=Requests( + { + ContainerDeclaration( + name="root_source", + container_type=RootDictContainer, + ): ObjectSet(MyKind.ROOT, {MyObjectID.root()}) + } + ), + analysis_configuration="", + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + assert isinstance(new_model, DictModel), "The analysis should return the same model type" + assert len(new_model) == 0, "PurgeAllAnalysis should empty the model" + stored_model = storage_provider.get_model() + assert stored_model == new_model.serialize(), "The model should be stored" + + +def test_schedule_serdes(model): + storage_provider = InMemoryStorageProvider() + + root1: ContainerDeclaration = ContainerDeclaration("root_source", RootDictContainer) + root2: ContainerDeclaration = ContainerDeclaration("root_destination", RootDictContainer) + child1: ContainerDeclaration = ContainerDeclaration("child_destination", ChildDictContainer) + child2: ContainerDeclaration = ContainerDeclaration("child_source", ChildDictContainer) + + declarations = [root1, root2, child1, child2] + + pipeline_configuration: PipelineConfiguration = {} + storage_provider.set_model(model.serialize()) + + begin_node = PipelineNode(SavePoint("begin", declarations)) + up_node = PipelineNode(ToHigherKindPipe(), bindings=[root1, child1]) + same_node = PipelineNode(SameKindPipe(), bindings=[child1, child2]) + down_node = PipelineNode(ToLowerKindPipe(), bindings=[child2, root2]) + end_node = PipelineNode(SavePoint("end", declarations)) + begin_node.add_successor(up_node).add_successor(same_node).add_successor( + down_node + ).add_successor(end_node) + + pipeline: Pipeline = Pipeline(set(declarations), begin_node) + assert begin_node.savepoint_range == SavePointsRange(start=1, end=2) + assert up_node.savepoint_range == SavePointsRange(start=1, end=2) + assert same_node.savepoint_range == SavePointsRange(start=1, end=2) + assert down_node.savepoint_range == SavePointsRange(start=1, end=2) + assert end_node.savepoint_range == SavePointsRange(start=2, end=2) + + root_obj = ObjectSet(MyKind.ROOT, {MyObjectID.root()}) + + # Force the savepoint to store the objects + container = RootDictContainer() + container.add_object(MyObjectID.root()) + requests = Requests({root1: root_obj}) + begin_node.run( + model=ReadOnlyModel(model), + containers={root1: container}, + incoming=requests, + outgoing=requests, + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + + schedule = pipeline.schedule( + model=ReadOnlyModel(model), + target_node=end_node, + requests=Requests({root2: root_obj}), + pipeline_configuration=pipeline_configuration, + storage_provider=storage_provider, + ) + schedule_str = schedule.serialize() + pipeline.deserialize_schedule(schedule_str)