缓存功能windows平台适配

This commit is contained in:
lujingze committed 2026-09-12 05:40:18 +00:00
1 parent 151e6e4b97
commit 22579e51c9
8 files changed
+333 -1

No files matched your search

+199
View File
@@ -0,0 +1,199 @@
"""Real native cache integration on both Windows and Linux.
The fixture uses the configured compiler and SUNDIALS installation directly.
CI sets SIMULATION_NATIVE_REQUIRE_TOOLCHAIN=1 so missing native dependencies
fail this suite instead of silently skipping platform acceptance.
"""
from __future__ import annotations
from dataclasses import replace
from hashlib import sha256
import json
import math
import os
from pathlib import Path
import subprocess
import tempfile
import threading
import unittest
from unittest.mock import patch
import xml.etree.ElementTree as ET
from app.main import compile_system_xml_network
from app.simulation.backends import simulation_config
from app.simulation.config import SolverActivityTracker
from app.simulation.native_codegen import build as native_build
from app.simulation.native_codegen.cache_storage import prune_cache
from app.simulation.native_codegen.compiler import compile_native_program
from app.simulation.native_codegen.runner import execute_native
from app.system_xml import validate_system_xml_document
ROOT = Path(__file__).resolve().parents[1]
FIXTURE = ROOT / "tests/fixtures/native-skill-test.xml"
class NativeCachePlatformTests(unittest.TestCase):
@classmethod
def setUpClass(cls):
try:
cls.compiler, cls.sundials, cls.compiler_version = native_build.toolchain()
except (OSError, RuntimeError, subprocess.SubprocessError) as exc:
message = f"Native cache platform test requires a working configured toolchain: {exc}"
if os.environ.get("SIMULATION_NATIVE_REQUIRE_TOOLCHAIN") == "1":
raise RuntimeError(message) from exc
raise unittest.SkipTest(message) from exc
def setUp(self):
self.temporary = tempfile.TemporaryDirectory(prefix="native-platform-")
self.addCleanup(self.temporary.cleanup)
# Exercise argument handling with spaces without assuming shell quoting.
self.root = Path(self.temporary.name) / "cache integration"
self.root.mkdir()
self.cache = self.root / "cache"
self.real_run = subprocess.run
self._builds = []
self.addCleanup(self._close_builds)
self.document = self._document(FIXTURE.read_bytes())
self.program = compile_native_program(compile_system_xml_network(self.document))
self.config = replace(simulation_config(self.document.simulation),
method="RK45", t_stop=0.1)
def _close_builds(self):
for build in reversed(self._builds):
build.close()
def _document(self, xml):
report = validate_system_xml_document(xml)
self.assertTrue(report.valid, report.as_dict())
self.assertIsNotNone(report.document)
return report.document
def _build(self, program):
commands = []
lock = threading.Lock()
def record(command, *args, **kwargs):
with lock:
commands.append(tuple(map(str, command)))
return self.real_run(command, *args, **kwargs)
with patch.object(native_build.subprocess, "run", side_effect=record):
build = native_build.build_native(program, cache_dir=self.cache)
self._builds.append(build)
return build, commands
@staticmethod
def _compile_count(commands):
return sum("-c" in command for command in commands)
@staticmethod
def _link_count(commands):
return sum("-o" in command and "-c" not in command and "-E" not in command
for command in commands)
def _execute(self, build, name, **kwargs):
payload = execute_native(build, self.config, 0.02,
run_dir=self.root / name, timeout=30, **kwargs)
self.assertTrue(payload["success"], payload.get("message"))
self.assertEqual(payload["simulatedUntil"], self.config.t_stop)
self.assertEqual(payload["series"]["time"][-1], self.config.t_stop)
self.assertEqual(set(payload["series"]), {"time", *(v.key for v in self.program.variables)})
self.assertTrue(all(math.isfinite(value)
for series in payload["series"].values() for value in series))
return payload
def _assert_windows_package(self, build):
if os.name != "nt":
return
self.assertEqual(build.executable.name, "model.exe")
names = [f"sundials_{name}.dll" for name in native_build.LIBRARIES]
if (self.sundials / "bin/vcruntime140.dll").is_file():
names.append("vcruntime140.dll")
for name in names:
with self.subTest(dll=name):
original = self.sundials / "bin" / name
packaged = build.executable.parent / name
self.assertTrue(packaged.is_file())
digest = sha256(packaged.read_bytes()).hexdigest()
self.assertEqual(digest, sha256(original.read_bytes()).hexdigest())
self.assertEqual(digest, build.manifest["artifacts"][name])
self.assertEqual(digest, build.manifest["buildIdentity"]["dependencies"][name])
# The EXE must load its copied DLLs without the compiler/Conda directories.
system32 = Path(os.environ.get("SystemRoot", r"C:\Windows")) / "System32"
completed = self.real_run(
[str(build.executable), "--init"], cwd=self.root,
env={**os.environ, "PATH": str(system32)},
capture_output=True, text=True, check=True, timeout=15,
)
self.assertEqual(len(json.loads(completed.stdout)), len(self.program.state_keys))
def test_real_build_execution_reuse_and_in_use_eviction(self):
self.assertFalse(self.cache.exists())
cold, cold_commands = self._build(self.program)
self.assertFalse(cold.cache_hit)
self.assertGreater(self._compile_count(cold_commands), 1)
self.assertEqual(self._link_count(cold_commands), 1)
self.assertEqual((cold.executable.parent / "model.c").read_bytes(), self.program.source.encode())
self.assertEqual((cold.executable.parent / "model.h").read_bytes(), self.program.header.encode())
self._assert_windows_package(cold)
cold_result = self._execute(cold, "cold run")
warm, warm_commands = self._build(self.program)
self.assertTrue(warm.cache_hit)
self.assertEqual(cold.manifest["buildKey"], warm.manifest["buildKey"])
self.assertEqual(self._compile_count(warm_commands), 0)
self.assertEqual(self._link_count(warm_commands), 0)
self.assertEqual(warm.details["objectCompilations"], 0)
self.assertEqual(warm.details["linkSeconds"], 0)
warm_result = self._execute(warm, "warm run")
self.assertEqual(cold_result["series"], warm_result["series"])
self.assertEqual(cold_result["finalState"], warm_result["finalState"])
xml = ET.fromstring(FIXTURE.read_bytes())
pressure = xml.find("./Components/Component[@id='amesim_pnch023_1']/Parameter[@name='p0']")
self.assertIsNotNone(pressure)
pressure.set("value", "16000000")
modified = self._document(ET.tostring(xml, encoding="utf-8"))
changed_program = compile_native_program(compile_system_xml_network(modified))
self.assertEqual(changed_program.header, self.program.header)
changed, changed_commands = self._build(changed_program)
self.assertFalse(changed.cache_hit)
self.assertNotEqual(changed.manifest["buildKey"], cold.manifest["buildKey"])
self.assertEqual(self._compile_count(changed_commands), 1)
self.assertEqual(self._link_count(changed_commands), 1)
self.assertEqual(changed.details["objectCompilations"], 1)
self.assertGreater(changed.details["objectCacheHits"], 0)
self._assert_windows_package(changed)
test = self
sweeps = []
class PruningTracker(SolverActivityTracker):
def start_integration(self, time):
super().start_integration(time)
# execute_native calls this after starting the real C child,
# before reading its result. Both model entries are still in use.
report = prune_cache(test.cache, model_limit_bytes=0)
sweeps.append(report)
test.assertGreaterEqual(report["models"]["skippedInUse"], 2)
test.assertTrue(cold.executable.is_file())
test.assertTrue(changed.executable.is_file())
changed_result = self._execute(changed, "changed run", activity_tracker=PruningTracker())
self.assertEqual(len(sweeps), 1)
column = "amesim_pnch023_1.p"
self.assertNotEqual(changed_result["series"][column][0], cold_result["series"][column][0])
self.assertEqual(changed_result["buildKey"], changed.manifest["buildKey"])
# Once use ends, a zero budget may evict the older model while retaining
# the final oversized entry. Release every shared handle explicitly.
self._close_builds()
report = prune_cache(self.cache, model_limit_bytes=0)
self.assertEqual(report["models"]["removedEntries"], 1)
self.assertFalse(cold.executable.exists())
self.assertTrue(changed.executable.is_file())
if __name__ == "__main__":
unittest.main()
+58
View File
@@ -6,6 +6,7 @@ import json
import os
from pathlib import Path
import subprocess
from types import SimpleNamespace
import sys
import tempfile
import unittest
@@ -108,6 +109,63 @@ with acquire_cache_lease(Path(sys.argv[1]), 'models', sys.argv[2], exclusive=Tru
gc.collect()
self.assertIsNotNone(self.lease(3, exclusive=True, blocking=False))
def test_touch_uses_available_utime_operation_and_preserves_lru(self):
real_utime = os.utime
for supported in (False, True):
with self.subTest(follow_symlinks_supported=supported):
self.cache = self.folder / f"cache-capability-{supported}"
touched = self.entry(0, 60, 1)
old = self.entry(1, 60, 2)
artifact_bytes = (touched / "artifact").read_bytes()
before = touched.stat().st_mtime_ns
calls = []
def platform_utime(path, times=None, **kwargs):
calls.append(kwargs)
if not supported and kwargs.get("follow_symlinks") is False:
raise NotImplementedError("utime: follow_symlinks unavailable on this platform")
# Both capability branches can run on either host. The
# real timestamp change is on a validated ordinary folder.
return real_utime(path, times)
with patch.object(storage.os, "utime", platform_utime), patch.object(
storage.os, "supports_follow_symlinks", {platform_utime} if supported else set(),
):
with self.lease(0):
storage.touch_cache_entry(self.cache, "models", key(0))
self.assertEqual(calls, [{"follow_symlinks": False}] if supported else [{}])
self.assertGreater(touched.stat().st_mtime_ns, before)
self.assertEqual((touched / "artifact").read_bytes(), artifact_bytes)
storage.prune_cache(self.cache, model_limit_bytes=60)
self.assertTrue(touched.exists())
self.assertFalse(old.exists())
def test_touch_rejects_windows_reparse_directory_before_timestamp_update(self):
entry = self.entry(0, 60, 1)
real_lstat = Path.lstat
attributes = SimpleNamespace(
st_mode=entry.lstat().st_mode, st_file_attributes=0x400,
)
def reparse_lstat(path, *args, **kwargs):
return attributes if path == entry else real_lstat(path, *args, **kwargs)
with patch.object(Path, "lstat", reparse_lstat), patch.object(
storage.os, "supports_follow_symlinks", set(),
), patch.object(storage.os, "utime") as update:
with self.assertRaisesRegex(RuntimeError, "real directory"):
storage.touch_cache_entry(self.cache, "models", key(0))
update.assert_not_called()
def test_touch_does_not_hide_timestamp_permission_errors(self):
self.entry(0, 60, 1)
with patch.object(storage.os, "supports_follow_symlinks", set()), patch.object(
storage.os, "utime", side_effect=PermissionError("timestamp denied"),
) as update:
with self.assertRaisesRegex(PermissionError, "timestamp denied"):
storage.touch_cache_entry(self.cache, "models", key(0))
update.assert_called_once()
def test_lru_removes_oldest_until_separate_budgets_are_met(self):
oldest = self.entry(0, 60, 1)
middle = self.entry(1, 60, 2)