404 lines
22 KiB
Python
404 lines
22 KiB
Python
"""On-demand native modules, checked object reuse and bounded executable caches."""
|
|
from __future__ import annotations
|
|
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
from dataclasses import dataclass, field
|
|
from hashlib import sha256
|
|
import json
|
|
import logging
|
|
import os
|
|
from pathlib import Path
|
|
import re
|
|
import shutil
|
|
import subprocess
|
|
import sys
|
|
import tempfile
|
|
import time
|
|
from typing import Any, Callable
|
|
|
|
from .compiler import NativeProgram
|
|
from .modules import component_modules
|
|
from .processes import command_environment, is_driver_launch_failure, is_transient_process_error, LAUNCH_RETRY_DELAYS
|
|
from .cache_storage import acquire_cache_lease, touch_cache_entry, prune_cache
|
|
|
|
LOGGER = logging.getLogger(__name__)
|
|
WINDOWS = os.name == 'nt'
|
|
|
|
|
|
class NativeCommandError(RuntimeError):
|
|
def __init__(self, message, command, *, exit_code=None, stdout=b'', stderr='', error_type='CalledProcessError', attempts=1):
|
|
super().__init__(message)
|
|
self.details = {'command': list(command), 'cwd': str(Path.cwd()), 'exitCode': exit_code,
|
|
'stdout': stdout.decode('utf-8',errors='replace')[-8000:], 'stderr': stderr[-8000:],
|
|
'errorType': error_type, 'attempts': attempts}
|
|
|
|
ROOT = Path(__file__).resolve().parents[3]
|
|
NATIVE = ROOT / "native"
|
|
CACHE = ROOT / "app/data/native-builds"
|
|
LIBRARIES = ("cvode", "core", "nvecserial", "sunmatrixdense", "sunlinsoldense")
|
|
COMPILER_FLAGS = ("-std=c11", "-O3", "-Wall", "-Wextra", "-Werror", "-ffp-contract=off", "-fno-fast-math")
|
|
RUNTIME_SOURCES = ("runtime/main.c", "runtime/common.c", "runtime/rk45.c",
|
|
"runtime/cvode_solver.c", "runtime/json_numbers.c", "encoding/ryu/d2s.c")
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class NativeBuild:
|
|
executable: Path
|
|
manifest: dict
|
|
cache_hit: bool
|
|
seconds: float
|
|
details: dict = field(default_factory=dict)
|
|
_lease: Any = field(default=None, repr=False, compare=False)
|
|
_cache: Path | None = field(default=None, repr=False, compare=False)
|
|
|
|
def close(self) -> None:
|
|
"""Release this caller's executable pin after its last use."""
|
|
if self._lease is not None:
|
|
self._lease.close()
|
|
if self._cache is not None:
|
|
_prune(self._cache)
|
|
|
|
|
|
def _hash(path: Path) -> str:
|
|
return sha256(path.read_bytes()).hexdigest()
|
|
|
|
|
|
def _key(value: object) -> str:
|
|
return sha256(json.dumps(value, sort_keys=True, separators=(",", ":")).encode()).hexdigest()
|
|
|
|
|
|
def _runtime_sources(program: NativeProgram) -> list[Path]:
|
|
return [NATIVE / name for name in RUNTIME_SOURCES] + [
|
|
NATIVE / "components/modules" / f"{name}.c"
|
|
for name in component_modules(program.source + "\n" + program.header)
|
|
]
|
|
|
|
|
|
def _budget(name: str, default: int) -> int:
|
|
value = os.environ.get(name, str(default))
|
|
if not re.fullmatch(r"[0-9]+", value):
|
|
raise ValueError(f"{name} must be a nonnegative integer in MiB.")
|
|
return int(value) * 1024**2
|
|
|
|
|
|
def _prune(cache: Path) -> dict:
|
|
return prune_cache(cache,
|
|
model_limit_bytes=_budget("SIMULATION_NATIVE_MODEL_CACHE_MB", 256),
|
|
object_limit_bytes=_budget("SIMULATION_NATIVE_OBJECT_CACHE_MB", 128))
|
|
|
|
|
|
def _checked_manifest(target: Path, key_name: str, key: str, version: int, required: set[str]) -> dict | None:
|
|
if not target.exists() and not target.is_symlink():
|
|
return None
|
|
message = f"Native cache integrity check failed: {target}"
|
|
if target.is_symlink() or not target.is_dir():
|
|
raise RuntimeError(message)
|
|
try:
|
|
path = target / "manifest.json"
|
|
if path.is_symlink():
|
|
raise ValueError("linked manifest")
|
|
manifest = json.loads(path.read_text(encoding="utf-8"))
|
|
if not isinstance(manifest, dict) or manifest.get(key_name) != key or manifest.get("cacheVersion") != version:
|
|
raise ValueError("identity mismatch")
|
|
artifacts = manifest.get("artifacts")
|
|
if not isinstance(artifacts, dict) or set(artifacts) != required:
|
|
raise ValueError("incomplete artifacts")
|
|
if key_name == "objectKey":
|
|
recipe = {name: manifest[name] for name in ("cacheVersion", "sourceName", "preprocessedSha256", "compiler")}
|
|
if _key(recipe) != key:
|
|
raise ValueError("object identity mismatch")
|
|
else:
|
|
recipe = manifest.get("buildIdentity")
|
|
if not isinstance(recipe, dict) or _key(recipe) != key or manifest.get("objectKeys") != recipe["objectKeys"]:
|
|
raise ValueError("model identity mismatch")
|
|
if artifacts["model.c"] != recipe["sourceSha256"] or artifacts["model.h"] != recipe["headerSha256"]:
|
|
raise ValueError("model source identity mismatch")
|
|
if _key({name: manifest[name] for name in recipe["contractKeys"]}) != recipe["contractSha256"]:
|
|
raise ValueError("model contract mismatch")
|
|
if artifacts["THIRD_PARTY_NOTICES.txt"] != recipe["notice"]:
|
|
raise ValueError("license artifact mismatch")
|
|
for name in required:
|
|
if name.lower().endswith(".dll") and artifacts[name] != recipe["dependencies"][name]:
|
|
raise ValueError("runtime dependency artifact mismatch")
|
|
for name, digest in artifacts.items():
|
|
if not re.fullmatch(r"[A-Za-z0-9_.-]+", name) or name in (".", ".."):
|
|
raise ValueError("invalid artifact name")
|
|
path = target / name
|
|
if path.is_symlink() or not path.is_file() or not isinstance(digest, str) or _hash(path) != digest:
|
|
raise ValueError("artifact mismatch")
|
|
return manifest
|
|
except (OSError, ValueError, TypeError, KeyError) as exc:
|
|
raise RuntimeError(message) from exc
|
|
|
|
|
|
def _publish(stage: Path, target: Path, key_name: str, key: str, version: int, required: set[str]) -> dict:
|
|
try:
|
|
stage.rename(target)
|
|
except OSError:
|
|
# Another reader/build may publish the same content while we compile.
|
|
# Its artifacts need not be byte-identical (e.g. PE timestamps), but
|
|
# must satisfy the same complete identity and its own recorded hashes.
|
|
winner = _checked_manifest(target, key_name, key, version, required)
|
|
if winner is None:
|
|
raise
|
|
shutil.rmtree(stage)
|
|
return winner
|
|
manifest = _checked_manifest(target, key_name, key, version, required)
|
|
assert manifest is not None
|
|
return manifest
|
|
|
|
|
|
def _command(command: list[str], *, log: list[str], timeout: float = 120,
|
|
require_stdout: bool = False) -> bytes:
|
|
environment = command_environment(command[0])
|
|
for attempt in range(len(LAUNCH_RETRY_DELAYS) + 1):
|
|
try:
|
|
result = subprocess.run(command, capture_output=True, timeout=timeout,
|
|
stdin=subprocess.DEVNULL, close_fds=True, env=environment,
|
|
creationflags=subprocess.CREATE_NO_WINDOW if WINDOWS else 0)
|
|
except OSError as exc:
|
|
if not is_transient_process_error(exc) or attempt == len(LAUNCH_RETRY_DELAYS):
|
|
raise
|
|
reason = f'{type(exc).__name__}: {exc}'
|
|
else:
|
|
stderr = result.stderr.decode('utf-8', errors='replace')
|
|
log.append(' '.join(command) + '\n' + stderr)
|
|
empty = result.returncode == 0 and require_stdout and not result.stdout.strip()
|
|
transient = is_driver_launch_failure(stderr) or (empty and WINDOWS)
|
|
if not result.returncode and not empty:
|
|
return result.stdout
|
|
reason = 'Compiler returned no identification output.' if empty else stderr[-3000:]
|
|
if not transient or attempt == len(LAUNCH_RETRY_DELAYS):
|
|
prefix = ('GCC 已启动,但编译子进程启动失败(已完成短间隔重试);'
|
|
if is_driver_launch_failure(stderr) else 'Native compilation failed: ')
|
|
raise NativeCommandError(prefix + reason, command, exit_code=result.returncode,
|
|
stdout=result.stdout, stderr=stderr, error_type='EmptyCompilerOutput' if empty else 'CalledProcessError',
|
|
attempts=attempt+1)
|
|
delay = LAUNCH_RETRY_DELAYS[attempt]
|
|
log.append(f'Compiler launch retry {attempt+1}/{len(LAUNCH_RETRY_DELAYS)} after {delay:g}s: {reason}')
|
|
LOGGER.warning('Compiler process launch failed; retry %d/%d in %.2fs: %s',
|
|
attempt+1, len(LAUNCH_RETRY_DELAYS), delay, reason.strip())
|
|
time.sleep(delay)
|
|
raise AssertionError('Unreachable compiler retry state')
|
|
|
|
|
|
def toolchain() -> tuple[str, Path, str]:
|
|
if os.name != "nt" and not sys.platform.startswith("linux"):
|
|
raise RuntimeError("Native builds support Windows and Linux; this platform is not supported.")
|
|
compiler = os.environ.get("SIMULATION_NATIVE_CC") or shutil.which("gcc")
|
|
if not compiler:
|
|
raise RuntimeError("C compiler not found; set SIMULATION_NATIVE_CC to gcc.")
|
|
configured = os.environ.get("SUNDIALS_ROOT")
|
|
candidates = ([Path(configured)] if configured else
|
|
[Path(sys.base_prefix) / "Library"] if os.name == "nt" else
|
|
[Path(sys.prefix) / "native/sundials-7.4.0", Path("/usr/local"), Path("/usr")])
|
|
base = next((path for path in candidates if (path / "include/cvode/cvode.h").is_file()), candidates[0])
|
|
if not (base / "include/cvode/cvode.h").is_file():
|
|
raise RuntimeError("SUNDIALS C development files not found; set SUNDIALS_ROOT.")
|
|
if os.name == 'nt':
|
|
compiler = str(Path(shutil.which(compiler) or compiler).resolve())
|
|
version = _command([compiler, '--version'], log=[], timeout=15, require_stdout=True).decode('utf-8', errors='replace').strip().splitlines()[0]
|
|
return compiler, base, version
|
|
|
|
|
|
def platform_build_inputs(sundials: Path) -> tuple[list[str], list[Path], list[Path], str]:
|
|
"""Share production flags and dependencies with the startup smoke check."""
|
|
flags = list(COMPILER_FLAGS)
|
|
executable_name = "model.exe" if os.name == "nt" else "model"
|
|
if os.name == "nt":
|
|
flags += ["-D__USE_MINGW_ANSI_STDIO=1", "-static-libgcc"]
|
|
libraries = [sundials / "lib" / f"sundials_{name}.lib" for name in LIBRARIES]
|
|
dlls = [sundials / "bin" / f"sundials_{name}.dll" for name in LIBRARIES]
|
|
if (sundials / "bin/vcruntime140.dll").is_file():
|
|
dlls.append(sundials / "bin/vcruntime140.dll")
|
|
else:
|
|
flags += ["-D_POSIX_C_SOURCE=200809L"]
|
|
candidates = [sundials / "lib", sundials / "lib64", sundials / "lib/x86_64-linux-gnu"]
|
|
directory = next((p for p in candidates if all((p / f"libsundials_{name}.a").is_file() for name in LIBRARIES)), None)
|
|
if directory is None:
|
|
raise RuntimeError("SUNDIALS static libraries not found; set SUNDIALS_ROOT.")
|
|
libraries = [directory / f"libsundials_{name}.a" for name in LIBRARIES]
|
|
dlls = []
|
|
return flags, libraries, dlls, executable_name
|
|
|
|
|
|
def link_library_arguments(libraries: list[Path]) -> list[str]:
|
|
arguments = list(map(str, libraries))
|
|
if sys.platform.startswith("linux"):
|
|
arguments = ["-Wl,--start-group", *arguments, "-Wl,--end-group"]
|
|
return arguments
|
|
|
|
|
|
def build_native(program: NativeProgram, *, cache_dir: Path | None = None,
|
|
progress_callback: Callable[[str], None] | None = None) -> NativeBuild:
|
|
started = time.perf_counter()
|
|
if progress_callback:
|
|
progress_callback("native-cache-check")
|
|
compiler, sundials, compiler_version = toolchain()
|
|
flags, libraries, dlls, executable_name = platform_build_inputs(sundials)
|
|
cache = (cache_dir or CACHE).resolve()
|
|
cache.mkdir(parents=True, exist_ok=True)
|
|
for kind in ("models", "objects"):
|
|
path = cache / kind
|
|
if path.is_symlink():
|
|
raise RuntimeError(f"Native cache directory must not be a symbolic link: {path}")
|
|
path.mkdir(exist_ok=True)
|
|
_budget("SIMULATION_NATIVE_MODEL_CACHE_MB", 256)
|
|
_budget("SIMULATION_NATIVE_OBJECT_CACHE_MB", 128)
|
|
work_key = sha256(os.urandom(32)).hexdigest()
|
|
work_lease = acquire_cache_lease(cache, "models", work_key)
|
|
try:
|
|
stage = Path(tempfile.mkdtemp(prefix=f"building-{work_key}-", dir=cache))
|
|
except BaseException:
|
|
work_lease.close()
|
|
raise
|
|
logs: list[str] = []
|
|
model_lease = None
|
|
try:
|
|
(stage / "model.c").write_text(program.source, encoding="utf-8", newline="\n")
|
|
(stage / "model.h").write_text(program.header, encoding="utf-8", newline="\n")
|
|
compiler_path = Path(shutil.which(compiler) or compiler).resolve()
|
|
compiler_identity = {"version": compiler_version, "binarySha256": _hash(compiler_path),
|
|
"target": _command([compiler, "-dumpmachine"], log=logs, require_stdout=True).decode().strip(),
|
|
"flags": flags, "platform": sys.platform}
|
|
runtime = _runtime_sources(program)
|
|
sources = [stage / "model.c", *runtime]
|
|
native_headers = {p: _hash(p) for p in (NATIVE / "include").rglob("*.h")}
|
|
dependency_headers = {p: _hash(p) for directory in ("sundials", "cvode", "nvector", "sunmatrix", "sunlinsol")
|
|
for p in (sundials / "include" / directory).glob("*.h")}
|
|
preprocessing_start = time.perf_counter()
|
|
# Preprocessed bytes include the actual transitive headers and all
|
|
# effective macros/ABI switches. Compile precisely these bytes, so a
|
|
# header edit cannot race the cache identity and the compiler input.
|
|
def preprocess(source: Path) -> dict:
|
|
local_log: list[str] = []
|
|
source_hash = _hash(source)
|
|
data = _command([compiler, *flags, "-E", "-P", f"-fmacro-prefix-map={stage}=/generated",
|
|
"-I", str(stage), "-I", str(NATIVE / "include"), "-I", str(sundials / "include"),
|
|
str(source)], log=local_log)
|
|
name = "model.c" if source == stage / "model.c" else "native/" + source.relative_to(NATIVE).as_posix()
|
|
if _hash(source) != source_hash:
|
|
raise RuntimeError(f"Native source changed during preprocessing: {source}")
|
|
digest = sha256(data).hexdigest()
|
|
identity = {"cacheVersion": 1, "sourceName": name, "preprocessedSha256": digest, "compiler": compiler_identity}
|
|
return {"sourceName": name, "sourceSha256": source_hash, "data": data,
|
|
"objectKey": _key(identity), "identity": identity, "log": local_log}
|
|
with ThreadPoolExecutor(max_workers=min(4, len(sources))) as pool:
|
|
units = list(pool.map(preprocess, sources))
|
|
for unit in units:
|
|
logs.extend(unit.pop("log"))
|
|
preprocessing_seconds = time.perf_counter() - preprocessing_start
|
|
if any(_hash(path) != digest for path, digest in {**native_headers, **dependency_headers}.items()):
|
|
raise RuntimeError("Native headers changed during preprocessing; retry with stable sources.")
|
|
dependencies = {p.name: _hash(p) for p in libraries + dlls}
|
|
notice = (NATIVE / "THIRD_PARTY_NOTICES.txt").read_bytes()
|
|
identity = {"cacheVersion": 2, "sourceSha256": sha256(program.source.encode()).hexdigest(),
|
|
"headerSha256": sha256(program.header.encode()).hexdigest(),
|
|
"contractSha256": _key(program.manifest()), "contractKeys": sorted(program.manifest()),
|
|
"objectKeys": [unit["objectKey"] for unit in units],
|
|
"compiler": compiler_identity, "dependencies": dependencies, "notice": sha256(notice).hexdigest()}
|
|
signature = _key(identity)
|
|
target = cache / "models" / signature
|
|
required = {"model.c", "model.h", executable_name, "THIRD_PARTY_NOTICES.txt", *(p.name for p in dlls)}
|
|
model_lease = acquire_cache_lease(cache, "models", signature)
|
|
manifest = _checked_manifest(target, "buildKey", signature, 2, required)
|
|
details = {"selectedModules": list(component_modules(program.source + "\n" + program.header)),
|
|
"unitCount": len(units), "objectCacheHits": 0, "objectCompilations": 0,
|
|
"preprocessSeconds": preprocessing_seconds, "compileSeconds": 0.0,
|
|
"compileWallSeconds": 0.0, "linkSeconds": 0.0, "modelCacheHit": manifest is not None}
|
|
details['compilerLaunchRetries'] = sum(line.startswith('Compiler launch retry ') for line in logs)
|
|
if manifest is not None:
|
|
if progress_callback:
|
|
progress_callback("native-cache-hit")
|
|
touch_cache_entry(cache, "models", signature)
|
|
details["pruning"] = _prune(cache)
|
|
result = NativeBuild(target / executable_name, manifest, True, time.perf_counter() - started,
|
|
details, model_lease, cache)
|
|
model_lease = None
|
|
return result
|
|
if progress_callback:
|
|
progress_callback("native-compilation")
|
|
compile_start = time.perf_counter()
|
|
def compile_unit(index_unit: tuple[int, dict]) -> dict:
|
|
index, unit = index_unit
|
|
key = unit["objectKey"]
|
|
destination = cache / "objects" / key
|
|
with acquire_cache_lease(cache, "objects", key):
|
|
object_manifest = _checked_manifest(destination, "objectKey", key, 1, {"unit.o"})
|
|
hit = object_manifest is not None
|
|
elapsed = 0.0
|
|
local_log: list[str] = []
|
|
if not hit:
|
|
temporary = Path(tempfile.mkdtemp(prefix=f"building-{key}-", dir=cache / "objects"))
|
|
try:
|
|
(temporary / "unit.i").write_bytes(unit["data"])
|
|
before = time.perf_counter()
|
|
_command([compiler, *flags, "-x", "cpp-output", "-c", str(temporary / "unit.i"),
|
|
"-o", str(temporary / "unit.o")], log=local_log)
|
|
elapsed = time.perf_counter() - before
|
|
(temporary / "unit.i").unlink()
|
|
object_manifest = {**unit["identity"], "objectKey": key, "artifacts": {"unit.o": _hash(temporary / "unit.o")}}
|
|
(temporary / "manifest.json").write_text(json.dumps(object_manifest, indent=2) + "\n")
|
|
_publish(temporary, destination, "objectKey", key, 1, {"unit.o"})
|
|
finally:
|
|
if temporary.exists():
|
|
shutil.rmtree(temporary)
|
|
# Linking uses private copies, so another model's LRU pass can
|
|
# safely reclaim the shared object after this lease is released.
|
|
local_object = stage / f"object-{index}.o"
|
|
shutil.copyfile(destination / "unit.o", local_object)
|
|
touch_cache_entry(cache, "objects", key)
|
|
return {"path": local_object, "hit": hit, "seconds": elapsed, "log": local_log}
|
|
with ThreadPoolExecutor(max_workers=min(4, len(units))) as pool:
|
|
objects = list(pool.map(compile_unit, enumerate(units)))
|
|
details["compileWallSeconds"] = time.perf_counter() - compile_start
|
|
details["compileSeconds"] = sum(item["seconds"] for item in objects)
|
|
details["objectCacheHits"] = sum(item["hit"] for item in objects)
|
|
details["objectCompilations"] = len(objects) - details["objectCacheHits"]
|
|
for item in objects:
|
|
logs.extend(item["log"])
|
|
link_libraries = link_library_arguments(libraries)
|
|
if progress_callback:
|
|
progress_callback("native-linking")
|
|
before_link = time.perf_counter()
|
|
_command([compiler, *flags, *[str(item["path"]) for item in objects], *link_libraries,
|
|
"-lm", "-o", str(stage / executable_name)], log=logs)
|
|
details['compilerLaunchRetries'] = sum(line.startswith('Compiler launch retry ') for line in logs)
|
|
details["linkSeconds"] = time.perf_counter() - before_link
|
|
if {p.name: _hash(p) for p in libraries + dlls} != dependencies:
|
|
raise RuntimeError("Native dependencies changed during linking; retry with a stable toolchain.")
|
|
for item in objects:
|
|
item["path"].unlink()
|
|
for library in dlls:
|
|
shutil.copyfile(library, stage / library.name)
|
|
if _hash(stage / library.name) != dependencies[library.name]:
|
|
raise RuntimeError("Native runtime dependency changed during copying; retry with a stable toolchain.")
|
|
(stage / "THIRD_PARTY_NOTICES.txt").write_bytes(notice)
|
|
(stage / "build.log").write_text("\n".join(logs), encoding="utf-8")
|
|
manifest = {**program.manifest(), "cacheVersion": 2, "buildKey": signature, "buildIdentity": identity,
|
|
"compiler": compiler_version, "compilerFlags": flags,
|
|
"sourceHashes": {unit["sourceName"]: unit["sourceSha256"] for unit in units},
|
|
"nativeHeaderHashes": {"native/" + p.relative_to(NATIVE).as_posix(): digest
|
|
for p, digest in native_headers.items()},
|
|
"dependencyHashes": dependencies, "selectedModules": details["selectedModules"],
|
|
"objectKeys": [unit["objectKey"] for unit in units],
|
|
"artifacts": {name: _hash(stage / name) for name in sorted(required)}}
|
|
(stage / "manifest.json").write_text(json.dumps(manifest, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
|
manifest = _publish(stage, target, "buildKey", signature, 2, required)
|
|
touch_cache_entry(cache, "models", signature)
|
|
details["pruning"] = _prune(cache)
|
|
result = NativeBuild(target / executable_name, manifest, False, time.perf_counter() - started,
|
|
details, model_lease, cache)
|
|
model_lease = None
|
|
return result
|
|
finally:
|
|
if model_lease is not None:
|
|
model_lease.close()
|
|
# Failed compilations never publish partial entries or retain unlimited
|
|
# model sources/preprocessed files. Error stderr is included in the exception.
|
|
try:
|
|
if stage.exists():
|
|
shutil.rmtree(stage)
|
|
finally:
|
|
work_lease.close()
|