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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 4 additions & 20 deletions activitysim/cli/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,13 +24,15 @@


INJECTABLES = [
"working_dir",
"data_dir",
"configs_dir",
"data_model_dir",
"output_dir",
"cache_dir",
"settings_file_name",
"imported_extensions",
"_extension_locations",
"run_timestamp",
"run_id",
]
Expand Down Expand Up @@ -163,31 +165,13 @@ def inject_arg(name, value):
if args.working_dir:
# activitysim will look in the current working directory for
# 'configs', 'data', and 'output' folders by default
args.working_dir = os.path.abspath(args.working_dir)
os.chdir(args.working_dir)

inject_arg("run_id", state.tracing.run_id)

if args.ext:
for e in args.ext:
basepath, extpath = os.path.split(e)
if not basepath:
basepath = "."
sys.path.insert(0, os.path.abspath(basepath))
try:
importlib.import_module(extpath)
except ImportError as err:
logger.exception("ImportError")
raise
except Exception as err:
logger.exception(f"Error {err}")
raise
finally:
del sys.path[0]
inject_arg("imported_extensions", args.ext)
else:
inject_arg("imported_extensions", ())

state.filesystem = FileSystem.parse_args(args)
state.import_extensions(args.ext or [], append=False)
for config_dir in state.filesystem.get_configs_dir():
if not config_dir.is_dir():
print(f"missing config directory: {config_dir}", file=sys.stderr)
Expand Down
181 changes: 181 additions & 0 deletions activitysim/cli/test/test_extensions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,181 @@
"""No-data, end-to-end regressions for CLI and Python API extension loading."""

from __future__ import annotations

import os
import subprocess
import sys

import pytest

EXTENSION = """\
import multiprocessing
from pathlib import Path
from activitysim.core import workflow
from .value import VALUE

@workflow.step(cache=True, kind="cached_object", overloading=True)
def network_los_preload(state: workflow.State):
return None

@workflow.step(cache=True, kind="cached_object", overloading=True)
def shadow_pricing_info(state: workflow.State):
return None

@workflow.step(cache=True, kind="cached_object", overloading=True)
def shadow_pricing_choice_info(state: workflow.State):
return None

@workflow.step
def extension_hello(state: workflow.State):
process = multiprocessing.current_process().name
Path(state.get_output_file_path("hello.txt")).write_text(f"{VALUE}:{process}")
"""

SETTINGS = """\
models: [extension_hello]
num_processes: 1
multiprocess_steps:
- name: mp_hello
begin: extension_hello
num_processes: 1
check_model_settings: true
memory_profile: false
sharrow: false
use_shadow_pricing: false
"""

API_RUNNER = """\
import importlib
import multiprocessing
import os
import sys
from pathlib import Path
from activitysim import abm
from activitysim.core import workflow

if __name__ == "__main__":
multiprocessing.set_start_method("spawn", force=True)
model, output, extension, multiprocess, elsewhere = sys.argv[1:]
state = workflow.State.make_default(Path(model), output_dir=Path(output))
state.import_extensions(extension)
# Existing example repositories use this public list as importable names.
for name in state.get_injectable("imported_extensions"):
checker = importlib.import_module(name + ".settings_checker")
assert checker.EXTENSION_CHECKER_SETTINGS == {}
state.settings.multiprocess = multiprocess == "yes"
os.chdir(elsewhere)
state.run.all()
"""


@pytest.fixture
def model(tmp_path):
root = tmp_path / "model space"
for directory in ("configs", "data", "extensions"):
(root / directory).mkdir(parents=True)
(root / "configs" / "settings.yaml").write_text(SETTINGS)
(root / "extensions" / "__init__.py").write_text(EXTENSION)
(root / "extensions" / "value.py").write_text("VALUE = 42\n")
(root / "extensions" / "settings_checker.py").write_text(
'print("EXTENSION_SETTINGS_CHECKER_IMPORTED", flush=True)\n'
"EXTENSION_CHECKER_SETTINGS = {}\n"
)
return root


def run_and_check(command, cwd, output, multiprocess):
env = os.environ.copy()
for variable in (
"MKL_NUM_THREADS",
"OMP_NUM_THREADS",
"OPENBLAS_NUM_THREADS",
"NUMBA_NUM_THREADS",
"VECLIB_MAXIMUM_THREADS",
"NUMEXPR_NUM_THREADS",
):
env[variable] = "1"
result = subprocess.run(
command,
cwd=cwd,
env=env,
text=True,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
timeout=120,
)
assert result.returncode == 0, result.stdout
assert (output / "hello.txt").read_text() == (
"42:mp_hello" if multiprocess else "42:MainProcess"
)
return result.stdout


@pytest.mark.parametrize("multiprocess", [False, True], ids=["single", "multiprocess"])
@pytest.mark.parametrize(
"form", ["relative", "absolute", "bare", "dot", "trailing", "working-dir"]
)
def test_cli_extensions(model, tmp_path, multiprocess, form):
output = tmp_path / "output"
cwd = tmp_path
extra = []
if form == "relative":
extension = os.path.join(model.name, "extensions")
elif form == "absolute":
extension = str(model / "extensions")
elif form == "working-dir":
# A relative -w must not be applied twice after chdir.
extension = "extensions"
extra = ["-w", model.name]
else:
cwd = model
extension = {
"bare": "extensions",
"dot": "./extensions",
"trailing": "extensions" + os.sep,
}[form]
command = [
sys.executable,
"-m",
"activitysim",
"run",
"-c",
str(model / "configs"),
"-d",
str(model / "data"),
"-o",
str(output),
"--ext",
extension,
*extra,
]
if multiprocess:
command += ["-m"]
stdout = run_and_check(command, cwd, output, multiprocess)
assert "EXTENSION_SETTINGS_CHECKER_IMPORTED" in stdout


@pytest.mark.parametrize("multiprocess", [False, True], ids=["single", "spawn"])
@pytest.mark.parametrize("form", ["relative", "absolute", "dot", "trailing"])
def test_api_extensions(model, tmp_path, multiprocess, form):
runner = tmp_path / "api_runner.py"
runner.write_text(API_RUNNER)
output = tmp_path / "output"
elsewhere = tmp_path / "unrelated cwd"
elsewhere.mkdir()
extension = {
"relative": "extensions",
"absolute": str(model / "extensions"),
"dot": "./extensions",
"trailing": "extensions" + os.sep,
}[form]
command = [
sys.executable,
str(runner),
str(model),
str(output),
extension,
"yes" if multiprocess else "no",
str(elsewhere),
]
run_and_check(command, tmp_path, output, multiprocess)
32 changes: 32 additions & 0 deletions activitysim/core/extensions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,32 @@
"""Shared extension loading for the CLI, workflow states, and workers."""

from __future__ import annotations

import importlib
import os
import sys
from pathlib import Path
from types import ModuleType


def resolve_extension(extension: str | os.PathLike, working_dir=None) -> str:
"""Freeze the search directory while retaining the Python module name.

The final path component is a module name (possibly dotted), not a Python
filename. Like Python imports, a module may also be found elsewhere on
sys.path. Store an absolute location so workers need not share the caller's
current directory. abspath normalizes separators and trailing slashes without
resolving symlinks, which could change the name of a package being imported.
"""
return os.path.abspath(Path(working_dir or Path.cwd()) / extension)


def import_extension(extension: str | os.PathLike) -> ModuleType:
"""Import a normalized extension location, restoring sys.path on failure too."""
location = Path(extension)
original_path = sys.path[:]
sys.path.insert(0, str(location.parent))
try:
return importlib.import_module(location.name)
finally:
sys.path[:] = original_path
15 changes: 2 additions & 13 deletions activitysim/core/mp_tasks.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@
from __future__ import annotations

import glob
import importlib
import logging
import multiprocessing
import os
Expand Down Expand Up @@ -928,18 +927,8 @@ def setup_injectables_and_logging(injectables, locutor: bool = True) -> workflow

# re-import extension modules to register injectables
ext = state.get_injectable("imported_extensions", default=())
for e in ext:
basepath, extpath = os.path.split(e)
if not basepath:
basepath = "."
sys.path.insert(0, basepath)
try:
importlib.import_module(e)
except ImportError as err:
logger.exception("ImportError")
raise
finally:
del sys.path[0]
locations = state.get_injectable("_extension_locations", default={})
state.import_extensions([locations.get(e, e) for e in ext], append=False)

state.add_injectable("is_sub_task", True)
state.add_injectable("locutor", locutor)
Expand Down
Loading
Loading